diff --git a/configure.ac b/configure.ac index ffb92d5e842..b973fa16c3c 100644 --- a/configure.ac +++ b/configure.ac @@ -5111,7 +5111,13 @@ then AM_CFLAGS="$AM_CFLAGS -DWOLFSSL_AESNI" if test "$KERNEL_MODE_DEFAULTS" = "yes" then - AM_CFLAGS="$AM_CFLAGS -DWC_C_DYNAMIC_FALLBACK" + # FIPS v7 and later pin one lane at build time; dev builds may + # still switch. Older validated modules keep what they were + # validated with. Tested numerically, and on the same condition + # settings.h uses, so a new v7.x or lean-* string cannot slip past. + AS_IF([test "${HAVE_FIPS_VERSION_MAJOR:-0}" -ge 7 && \ + test "x$ENABLED_FIPS_DEV" != "xyes"],[], + [AM_CFLAGS="$AM_CFLAGS -DWC_C_DYNAMIC_FALLBACK"]) fi if test "$CC" != "icc" then @@ -14223,6 +14229,14 @@ AM_CFLAGS="$AM_CFLAGS $EXTRA_CFLAGS" AM_CCASFLAGS="$AM_CCASFLAGS $EXTRA_CCASFLAGS" AM_LDFLAGS="$AM_LDFLAGS $EXTRA_LDFLAGS" +# FIPS v7 and later never switch lanes at run time, however the define +# arrives. Same condition as the auto-add above and as settings.h. +AS_IF([test "${HAVE_FIPS_VERSION_MAJOR:-0}" -ge 7 && \ + test "x$ENABLED_FIPS_DEV" != "xyes"], + [AS_CASE([" $AM_CPPFLAGS $AM_CFLAGS $CPPFLAGS $CFLAGS "], + [*-DWC_C_DYNAMIC_FALLBACK*], + [AC_MSG_ERROR([WC_C_DYNAMIC_FALLBACK is not allowed with --enable-fips=$FIPS_VERSION; use --enable-fips=dev or --enable-fips=dev-no-post])])]) + CREATE_HEX_VERSION AC_SUBST([AM_CPPFLAGS]) AC_SUBST([AM_CFLAGS]) diff --git a/linuxkm/lkcapi_aes_glue.c b/linuxkm/lkcapi_aes_glue.c index c7b80f0ba40..eb471701ca7 100644 --- a/linuxkm/lkcapi_aes_glue.c +++ b/linuxkm/lkcapi_aes_glue.c @@ -2235,8 +2235,14 @@ static int ccmAesAead_rfc4309_loaded = 0; #error LKCAPI registration of AES-XTS requires WOLFSSL_AESXTS_STREAM (--enable-aesxts-stream). #endif -#if defined(WOLFSSL_AESNI) && !defined(WC_C_DYNAMIC_FALLBACK) && !defined(WC_DEBUG_FORCE_KERNEL_SETTINGS) - #error LKCAPI registration of AES-XTS with AESNI requires WC_C_DYNAMIC_FALLBACK. +/* AES-XTS asm needs a vector save on every call. Without the fallback the + * whole context is pinned to C at setkey, so no save is ever taken. */ +#if defined(WOLFSSL_AESNI) && !defined(WC_C_DYNAMIC_FALLBACK) && \ + !defined(WC_DEBUG_FORCE_KERNEL_SETTINGS) + #define WC_LINUXKM_XTS_NO_AESNI + #ifndef WC_FLAG_DONT_USE_VECTOR_OPS + #error AES-XTS without WC_C_DYNAMIC_FALLBACK needs WC_FLAG_DONT_USE_VECTOR_OPS. + #endif #endif struct km_AesXtsCtx { @@ -2286,6 +2292,15 @@ static int km_AesXtsSetKey(struct crypto_skcipher *tfm, const u8 *in_key, int err; struct km_AesXtsCtx * ctx = crypto_skcipher_ctx(tfm); +#ifdef WC_LINUXKM_XTS_NO_AESNI + /* Set before the key schedule is built so the C schedule is the one made. */ + ctx->aesXts->aes.use_aesni = WC_FLAG_DONT_USE_VECTOR_OPS; + ctx->aesXts->tweak.use_aesni = WC_FLAG_DONT_USE_VECTOR_OPS; +#ifdef WC_AES_XTS_SUPPORT_SIMULTANEOUS_ENC_AND_DEC_KEYS + ctx->aesXts->aes_decrypt.use_aesni = WC_FLAG_DONT_USE_VECTOR_OPS; +#endif +#endif + err = wc_AesXtsSetKeyNoInit(ctx->aesXts, in_key, key_len, AES_ENCRYPTION_AND_DECRYPTION); @@ -2296,12 +2311,6 @@ static int km_AesXtsSetKey(struct crypto_skcipher *tfm, const u8 *in_key, return -EINVAL; } - /* It's possible to set ctx->aesXts->{tweak,aes,aes_decrypt}.use_aesni to - * WC_FLAG_DONT_USE_VECTOR_OPS here, for WC_LINUXKM_C_FALLBACK_IN_SHIMS in - * AES-XTS, but we can use the WC_C_DYNAMIC_FALLBACK mechanism - * unconditionally because there's no AES-XTS in Cert 4718. - */ - #ifdef WOLFKM_DEBUG_AES pr_info("info: exiting km_AesXtsSetKey: %d\n", key_len); #endif /* WOLFKM_DEBUG_AES */ @@ -2321,7 +2330,8 @@ static int km_AesXtsSetKey(struct crypto_skcipher *tfm, const u8 *in_key, typeof(wc_AesXtsEncryptUpdate_fips) wc_AesXtsEncryptUpdate; #endif -#if defined(WOLFSSL_USE_SAVE_VECTOR_REGISTERS) && !defined(WC_LINUXKM_SVR_NO_BATCHING) +#if defined(WOLFSSL_USE_SAVE_VECTOR_REGISTERS) && \ + !defined(WC_LINUXKM_SVR_NO_BATCHING) && !defined(WC_LINUXKM_XTS_NO_AESNI) #ifndef WC_LINUXKM_XTS_SVR_BATCH #define WC_LINUXKM_XTS_SVR_BATCH (16 * 4096) #endif diff --git a/linuxkm/lkcapi_glue.c b/linuxkm/lkcapi_glue.c index 42d26a251ad..0692eb355eb 100644 --- a/linuxkm/lkcapi_glue.c +++ b/linuxkm/lkcapi_glue.c @@ -97,7 +97,13 @@ #define LKCAPI_HAVE_ARCH_ACCEL #endif -#if defined(LKCAPI_HAVE_ARCH_ACCEL) && \ +/* v7 pins one lane per algorithm, so the shims keep no second key schedule. + * Tried and failed to make a refused save reach one: skcipher from hardirq is + * refused (crypto/skcipher.c:449), softirq always has SIMD (fpu/core.c:76). */ +#if defined(HAVE_FIPS) && FIPS_VERSION3_GE(7,0,0) && \ + !defined(WOLFSSL_FIPS_DEV) && !defined(WOLFSSL_FIPS_DEV_NO_POST) + #undef WC_LINUXKM_C_FALLBACK_IN_SHIMS +#elif defined(LKCAPI_HAVE_ARCH_ACCEL) && \ (!defined(WC_C_DYNAMIC_FALLBACK) || \ (defined(HAVE_FIPS) && FIPS_VERSION3_LT(6,0,0))) && \ !defined(WC_LINUXKM_C_FALLBACK_IN_SHIMS) @@ -106,6 +112,16 @@ #undef WC_LINUXKM_C_FALLBACK_IN_SHIMS #endif +/* With one lane and no fallback, the save has to be available in every context + * the kernel may call us from. The module's own XSAVE/FXSAVE area supplies it + * in hardirq and NMI as well (linuxkm/x86_vector_register_glue.c). */ +#if defined(HAVE_FIPS) && FIPS_VERSION3_GE(7,0,0) && \ + !defined(WOLFSSL_FIPS_DEV) && !defined(WOLFSSL_FIPS_DEV_NO_POST) && \ + defined(CONFIG_X86) && defined(WOLFSSL_USE_SAVE_VECTOR_REGISTERS) && \ + !defined(WC_SVR_USE_NATIVE_REG_BUFS) + #error FIPS v7 LKCAPI needs WC_SVR_USE_NATIVE_REG_BUFS: one lane, no fallback. +#endif + #if defined(WC_LINUXKM_C_FALLBACK_IN_SHIMS) && !defined(CAN_SAVE_VECTOR_REGISTERS) #error WC_LINUXKM_C_FALLBACK_IN_SHIMS is defined but CAN_SAVE_VECTOR_REGISTERS is missing. #endif diff --git a/linuxkm/module_hooks.c b/linuxkm/module_hooks.c index e1807d28a55..b448041f8ee 100644 --- a/linuxkm/module_hooks.c +++ b/linuxkm/module_hooks.c @@ -1262,61 +1262,13 @@ static int wolfssl_init(void) #ifdef WC_LINUXKM_SVR_DYNAMIC_AUDITING { long long unsigned int svr_disallowed_count = wc_svr_disallowed_count_current(); - long long unsigned int svr_disallowed_snapshot; if (svr_disallowed_count > 0) { pr_err("ERROR: wc_svr_disallowed_count_current() returned %llu after wc_RunAllCast_fips().\n", svr_disallowed_count); (void)libwolfssl_cleanup(); return -ECANCELED; } - - #ifdef WC_LINUXKM_HAVE_STACK_DEBUG - { - unsigned long stack_usage; - wc_linuxkm_stack_hwm_prepare(0xee); - #endif - - ret = DISABLE_VECTOR_REGISTERS(); - if (ret != 0) { - pr_err("ERROR: DISABLE_VECTOR_REGISTERS() for wc_RunAllCast_fips() returned %d.\n", ret); - (void)libwolfssl_cleanup(); - return -ECANCELED; - } - - /* See the snapshot rationale in the wolfCrypt_IntegrityTest_fips() - * block above. */ - svr_disallowed_snapshot = wc_svr_disallowed_count_current(); - - ret = wc_RunAllCast_fips(); - - REENABLE_VECTOR_REGISTERS(); - - #ifdef WC_LINUXKM_HAVE_STACK_DEBUG - stack_usage = wc_linuxkm_stack_hwm_measure_rel(0xee); - pr_info("STACK INFO: rel usage by wc_RunAllCast_fips() with DISABLE_VECTOR_REGISTERS(): %lu\n", stack_usage); - /* shush up false stack HWM reading by kernel: */ - wc_linuxkm_stack_hwm_prepare(0); - } - #endif - - svr_disallowed_count = wc_svr_disallowed_count_current(); - if (svr_disallowed_count <= svr_disallowed_snapshot) { - pr_err("ERROR: wc_svr_disallowed_count_current() returned %llu after wc_RunAllCast_fips() with DISABLE_VECTOR_REGISTERS() (snapshot %llu): inhibited-save instrumentation was not exercised.\n", svr_disallowed_count, svr_disallowed_snapshot); - (void)libwolfssl_cleanup(); - return -ECANCELED; - } - - if (ret != 0) { - pr_err("ERROR: wc_RunAllCast_fips() with DISABLE_VECTOR_REGISTERS() returned %d.\n", ret); - (void)libwolfssl_cleanup(); - return -ECANCELED; - } - - ret = wolfCrypt_GetStatus_fips(); - if (ret != 0) { - pr_err("ERROR: wolfCrypt_GetStatus_fips() failed with code %d: %s\n", ret, wc_GetErrorString(ret)); - (void)libwolfssl_cleanup(); - return -ECANCELED; - } + /* The CASTs are not re-run with the registers disabled: CPUID picks + * one lane per algorithm, so a refused save is an error there. */ } #endif /* WC_LINUXKM_SVR_DYNAMIC_AUDITING */ diff --git a/linuxkm/x86_vector_register_glue.c b/linuxkm/x86_vector_register_glue.c index 7b1a34f5998..b4c66d74cec 100644 --- a/linuxkm/x86_vector_register_glue.c +++ b/linuxkm/x86_vector_register_glue.c @@ -599,8 +599,8 @@ WARN_UNUSED_RESULT int wc_save_vector_registers_x86(enum wc_svr_flags flags) * Note that this is not actually an abnormal condition -- e.g. with * LINUXKM_DRBG_GET_RANDOM_BYTES, get_random_u32() and the like called from * hard IRQ handlers can land here, and we return success if - * WC_SVR_USE_NATIVE_REG_BUFS, else WC_ACCEL_INHIBIT_E for graceful fallback - * to C. + * WC_SVR_USE_NATIVE_REG_BUFS, else WC_ACCEL_INHIBIT_E. Callers whose lane + * is pinned report that error rather than computing the answer in C. */ if ((cur_preempt_count & (NMI_MASK | HARDIRQ_MASK)) != 0) { #ifdef WC_SVR_USE_NATIVE_REG_BUFS diff --git a/tests/api/test_mldsa.c b/tests/api/test_mldsa.c index f3749d684b4..f7df6adc73e 100644 --- a/tests/api/test_mldsa.c +++ b/tests/api/test_mldsa.c @@ -30073,10 +30073,10 @@ int test_mldsa_encode_w1_large_values(void) #if defined(DEBUG_VECTOR_REGISTER_ACCESS) && \ defined(DEBUG_VECTOR_REGISTER_ACCESS_FUZZING) - /* Pin dispatch to the C path: under SVR2 fuzzing the two calls can - * otherwise take different (AVX2 vs C) implementations, which are only - * specified - and only equal - on the valid input domain. */ - WC_DEBUG_SET_VECTOR_REGISTERS_RETVAL(WC_NO_ERR_TRACE(SYSLIB_FAILED_E)); + /* Let every save succeed for this test. A refused save is an error that + * writes nothing, so two refused calls would compare equal while encoding + * nothing; pinning saves ON keeps one lane for both calls instead. */ + WC_DEBUG_SET_VECTOR_REGISTERS_RETVAL(0); #endif /* ---- 6-bit encoding (mldsa_encode_w1_88 path) ---- */ @@ -30093,8 +30093,8 @@ int test_mldsa_encode_w1_large_values(void) XMEMSET(enc_a, 0, sizeof(enc_a)); XMEMSET(enc_b, 0, sizeof(enc_b)); - wc_mldsa_encode_w1_88(w1, enc_a); - wc_mldsa_encode_w1_88(w1, enc_b); + ExpectIntEQ(wc_mldsa_encode_w1_88(w1, enc_a), 0); + ExpectIntEQ(wc_mldsa_encode_w1_88(w1, enc_b), 0); /* Determinism: same input must produce same output */ ExpectIntEQ(XMEMCMP(enc_a, enc_b, sizeof(enc_a)), 0); @@ -30106,8 +30106,8 @@ int test_mldsa_encode_w1_large_values(void) } XMEMSET(enc_a, 0, sizeof(enc_a)); XMEMSET(enc_b, 0, sizeof(enc_b)); - wc_mldsa_encode_w1_88(w1, enc_a); - wc_mldsa_encode_w1_88(w1, enc_b); + ExpectIntEQ(wc_mldsa_encode_w1_88(w1, enc_a), 0); + ExpectIntEQ(wc_mldsa_encode_w1_88(w1, enc_b), 0); ExpectIntEQ(XMEMCMP(enc_a, enc_b, sizeof(enc_a)), 0); /* Ascending pattern: each element differs */ @@ -30116,8 +30116,8 @@ int test_mldsa_encode_w1_large_values(void) } XMEMSET(enc_a, 0, sizeof(enc_a)); XMEMSET(enc_b, 0, sizeof(enc_b)); - wc_mldsa_encode_w1_88(w1, enc_a); - wc_mldsa_encode_w1_88(w1, enc_b); + ExpectIntEQ(wc_mldsa_encode_w1_88(w1, enc_a), 0); + ExpectIntEQ(wc_mldsa_encode_w1_88(w1, enc_b), 0); ExpectIntEQ(XMEMCMP(enc_a, enc_b, sizeof(enc_a)), 0); } #endif /* !WOLFSSL_NO_ML_DSA_44 */ @@ -30136,8 +30136,8 @@ int test_mldsa_encode_w1_large_values(void) XMEMSET(enc_a, 0, sizeof(enc_a)); XMEMSET(enc_b, 0, sizeof(enc_b)); - wc_mldsa_encode_w1_32(w1, enc_a); - wc_mldsa_encode_w1_32(w1, enc_b); + ExpectIntEQ(wc_mldsa_encode_w1_32(w1, enc_a), 0); + ExpectIntEQ(wc_mldsa_encode_w1_32(w1, enc_b), 0); ExpectIntEQ(XMEMCMP(enc_a, enc_b, sizeof(enc_a)), 0); } @@ -30148,8 +30148,8 @@ int test_mldsa_encode_w1_large_values(void) } XMEMSET(enc_a, 0, sizeof(enc_a)); XMEMSET(enc_b, 0, sizeof(enc_b)); - wc_mldsa_encode_w1_32(w1, enc_a); - wc_mldsa_encode_w1_32(w1, enc_b); + ExpectIntEQ(wc_mldsa_encode_w1_32(w1, enc_a), 0); + ExpectIntEQ(wc_mldsa_encode_w1_32(w1, enc_b), 0); ExpectIntEQ(XMEMCMP(enc_a, enc_b, sizeof(enc_a)), 0); /* Ascending pattern */ @@ -30158,8 +30158,8 @@ int test_mldsa_encode_w1_large_values(void) } XMEMSET(enc_a, 0, sizeof(enc_a)); XMEMSET(enc_b, 0, sizeof(enc_b)); - wc_mldsa_encode_w1_32(w1, enc_a); - wc_mldsa_encode_w1_32(w1, enc_b); + ExpectIntEQ(wc_mldsa_encode_w1_32(w1, enc_a), 0); + ExpectIntEQ(wc_mldsa_encode_w1_32(w1, enc_b), 0); ExpectIntEQ(XMEMCMP(enc_a, enc_b, sizeof(enc_a)), 0); } #endif /* !WOLFSSL_NO_ML_DSA_65 || !WOLFSSL_NO_ML_DSA_87 */ diff --git a/tests/api/test_mlkem.c b/tests/api/test_mlkem.c index 7f06dfc5564..11929bbe685 100644 --- a/tests/api/test_mlkem.c +++ b/tests/api/test_mlkem.c @@ -5037,3 +5037,296 @@ int test_wc_mlkem_cb_pending_rejected(void) #endif return EXPECT_RESULT(); } + +/* A refused save while re-decoding a public key must leave the key unusable, + * not holding one key's polynomials with another key's seed. */ +int test_wc_mlkem_decode_pubkey_refused_save(void) +{ + EXPECT_DECLS; +#if defined(WOLFSSL_HAVE_MLKEM) && !defined(WOLFSSL_NO_ML_KEM) && \ + !defined(WOLFSSL_MLKEM_NO_MAKE_KEY) && \ + !defined(WOLFSSL_MLKEM_NO_ENCAPSULATE) && \ + defined(DEBUG_VECTOR_REGISTER_ACCESS) + MlKemKey* kA = NULL; + MlKemKey* kB = NULL; + MlKemKey* k = NULL; + WC_RNG rng; + byte pkA[WC_ML_KEM_MAX_PUBLIC_KEY_SIZE]; + byte pkB[WC_ML_KEM_MAX_PUBLIC_KEY_SIZE]; + byte ct[WC_ML_KEM_MAX_CIPHER_TEXT_SIZE]; + byte ss[WC_ML_KEM_SS_SZ]; + word32 pubLen = 0; + int ret = 0; +#ifndef WOLFSSL_NO_ML_KEM_768 + const int mlkemType = WC_ML_KEM_768; +#elif !defined(WOLFSSL_NO_ML_KEM_512) + const int mlkemType = WC_ML_KEM_512; +#else + const int mlkemType = WC_ML_KEM_1024; +#endif + + XMEMSET(&rng, 0, sizeof(rng)); + ExpectIntEQ(wc_InitRng(&rng), 0); + ExpectNotNull(kA = (MlKemKey*)XMALLOC(sizeof(*kA), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (kA != NULL) { + XMEMSET(kA, 0, sizeof(*kA)); + } + ExpectNotNull(kB = (MlKemKey*)XMALLOC(sizeof(*kB), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (kB != NULL) { + XMEMSET(kB, 0, sizeof(*kB)); + } + ExpectNotNull(k = (MlKemKey*)XMALLOC(sizeof(*k), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (k != NULL) { + XMEMSET(k, 0, sizeof(*k)); + } + ExpectIntEQ(wc_MlKemKey_Init(kA, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_Init(kB, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_Init(k, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_MakeKey(kA, &rng), 0); + ExpectIntEQ(wc_MlKemKey_MakeKey(kB, &rng), 0); + ExpectIntEQ(wc_MlKemKey_PublicKeySize(kA, &pubLen), 0); + ExpectIntEQ(wc_MlKemKey_EncodePublicKey(kA, pkA, pubLen), 0); + ExpectIntEQ(wc_MlKemKey_EncodePublicKey(kB, pkB, pubLen), 0); + ExpectIntEQ(wc_MlKemKey_DecodePublicKey(k, pkA, pubLen), 0); + + if (EXPECT_SUCCESS()) { + WC_DEBUG_SET_VECTOR_REGISTERS_RETVAL( + WC_NO_ERR_TRACE(WC_ACCEL_INHIBIT_E)); + ret = wc_MlKemKey_DecodePublicKey(k, pkB, pubLen); + WC_DEBUG_SET_VECTOR_REGISTERS_RETVAL(0); + } + /* Only a build whose decode takes a save is refused here. */ + if (ret != 0) { + ExpectIntNE(wc_MlKemKey_Encapsulate(k, ct, ss, &rng), 0); + } + + wc_MlKemKey_Free(k); + wc_MlKemKey_Free(kB); + wc_MlKemKey_Free(kA); + XFREE(k, NULL, DYNAMIC_TYPE_TMP_BUFFER); + XFREE(kB, NULL, DYNAMIC_TYPE_TMP_BUFFER); + XFREE(kA, NULL, DYNAMIC_TYPE_TMP_BUFFER); + DoExpectIntEQ(wc_FreeRng(&rng), 0); +#endif + return EXPECT_RESULT(); +} + +/* A key reused for a decoded key must use the new key's matrix, not one cached + * from the key it held before. */ +int test_wc_mlkem_decode_reused_key(void) +{ + EXPECT_DECLS; +#if defined(WOLFSSL_HAVE_MLKEM) && !defined(WOLFSSL_NO_ML_KEM) && \ + !defined(WOLFSSL_MLKEM_NO_MAKE_KEY) && \ + !defined(WOLFSSL_MLKEM_NO_ENCAPSULATE) && \ + !defined(WOLFSSL_MLKEM_NO_DECAPSULATE) + MlKemKey* kA = NULL; + MlKemKey* kB = NULL; + WC_RNG rng; + byte pkB[WC_ML_KEM_MAX_PUBLIC_KEY_SIZE]; + byte skB[WC_ML_KEM_MAX_PRIVATE_KEY_SIZE]; + byte ct[WC_ML_KEM_MAX_CIPHER_TEXT_SIZE]; + byte ss1[WC_ML_KEM_SS_SZ]; + byte ss2[WC_ML_KEM_SS_SZ]; + word32 pubLen = 0; + word32 privLen = 0; + word32 ctLen = 0; +#ifndef WOLFSSL_NO_ML_KEM_768 + const int mlkemType = WC_ML_KEM_768; +#elif !defined(WOLFSSL_NO_ML_KEM_512) + const int mlkemType = WC_ML_KEM_512; +#else + const int mlkemType = WC_ML_KEM_1024; +#endif + + XMEMSET(&rng, 0, sizeof(rng)); + ExpectIntEQ(wc_InitRng(&rng), 0); + ExpectNotNull(kA = (MlKemKey*)XMALLOC(sizeof(*kA), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (kA != NULL) { + XMEMSET(kA, 0, sizeof(*kA)); + } + ExpectNotNull(kB = (MlKemKey*)XMALLOC(sizeof(*kB), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (kB != NULL) { + XMEMSET(kB, 0, sizeof(*kB)); + } + ExpectIntEQ(wc_MlKemKey_Init(kA, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_Init(kB, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_MakeKey(kB, &rng), 0); + ExpectIntEQ(wc_MlKemKey_PublicKeySize(kB, &pubLen), 0); + ExpectIntEQ(wc_MlKemKey_PrivateKeySize(kB, &privLen), 0); + ExpectIntEQ(wc_MlKemKey_CipherTextSize(kB, &ctLen), 0); + ExpectIntEQ(wc_MlKemKey_EncodePublicKey(kB, pkB, pubLen), 0); + ExpectIntEQ(wc_MlKemKey_EncodePrivateKey(kB, skB, privLen), 0); + + /* Public key into a made key: B must recover what A encapsulated. */ + ExpectIntEQ(wc_MlKemKey_MakeKey(kA, &rng), 0); + ExpectIntEQ(wc_MlKemKey_DecodePublicKey(kA, pkB, pubLen), 0); + ExpectIntEQ(wc_MlKemKey_Encapsulate(kA, ct, ss1, &rng), 0); + ExpectIntEQ(wc_MlKemKey_Decapsulate(kB, ss2, ct, ctLen), 0); + ExpectBufEQ(ss1, ss2, sizeof(ss1)); + + /* Private key into a made key: A must recover what B encapsulated. */ + ExpectIntEQ(wc_MlKemKey_MakeKey(kA, &rng), 0); + ExpectIntEQ(wc_MlKemKey_DecodePrivateKey(kA, skB, privLen), 0); + ExpectIntEQ(wc_MlKemKey_Encapsulate(kB, ct, ss1, &rng), 0); + ExpectIntEQ(wc_MlKemKey_Decapsulate(kA, ss2, ct, ctLen), 0); + ExpectBufEQ(ss1, ss2, sizeof(ss1)); + + /* A failed public key decode into a full key: decapsulate is refused. */ + ExpectIntEQ(wc_MlKemKey_MakeKey(kA, &rng), 0); + ExpectIntEQ(wc_MlKemKey_Encapsulate(kA, ct, ss1, &rng), 0); + pkB[0] = 0xff; + pkB[1] = 0xff; + ExpectIntNE(wc_MlKemKey_DecodePublicKey(kA, pkB, pubLen), 0); + ExpectIntEQ(wc_MlKemKey_Decapsulate(kA, ss2, ct, ctLen), + WC_NO_ERR_TRACE(BAD_STATE_E)); + + wc_MlKemKey_Free(kB); + wc_MlKemKey_Free(kA); + XFREE(kB, NULL, DYNAMIC_TYPE_TMP_BUFFER); + XFREE(kA, NULL, DYNAMIC_TYPE_TMP_BUFFER); + DoExpectIntEQ(wc_FreeRng(&rng), 0); +#endif + return EXPECT_RESULT(); +} + +#if defined(WOLFSSL_HAVE_MLKEM) && !defined(WOLFSSL_NO_ML_KEM) && \ + defined(WOLFSSL_MLKEM_DYNAMIC_KEYS) && defined(USE_WOLFSSL_MEMORY) && \ + !defined(WOLFSSL_STATIC_MEMORY) && !defined(WOLFSSL_DEBUG_MEMORY) && \ + !defined(WOLFSSL_MLKEM_NO_MAKE_KEY) && \ + !defined(WOLFSSL_MLKEM_NO_ENCAPSULATE) && \ + !defined(WOLFSSL_MLKEM_NO_DECAPSULATE) +#define MLKEM_OOM_TEST + +static wolfSSL_Malloc_cb mlkem_oom_mf; +static wolfSSL_Free_cb mlkem_oom_ff; +static wolfSSL_Realloc_cb mlkem_oom_rf; +static size_t mlkem_oom_size; +static int mlkem_oom_armed; + +/* Fails the next allocation of one key vector, then behaves normally. */ +static void* mlkem_oom_malloc(size_t n) +{ + if (mlkem_oom_armed && (n == mlkem_oom_size)) { + mlkem_oom_armed = 0; + return NULL; + } + return (mlkem_oom_mf != NULL) ? mlkem_oom_mf(n) : malloc(n); +} + +static void mlkem_oom_free(void* p) +{ + if (mlkem_oom_ff != NULL) { + mlkem_oom_ff(p); + } + else { + free(p); + } +} + +static void* mlkem_oom_realloc(void* p, size_t n) +{ + return (mlkem_oom_rf != NULL) ? mlkem_oom_rf(p, n) : realloc(p, n); +} + +/* Decodes into kA with one key vector allocation failing. */ +static int mlkem_oom_decode(MlKemKey* kA, int priv, const byte* in, + word32 inLen) +{ + int ret; + + if ((wolfSSL_GetAllocators(&mlkem_oom_mf, &mlkem_oom_ff, + &mlkem_oom_rf) != 0) || + (wolfSSL_SetAllocators(mlkem_oom_malloc, mlkem_oom_free, + mlkem_oom_realloc) != 0)) { + return -1; + } + mlkem_oom_armed = 1; + if (priv) { + ret = wc_MlKemKey_DecodePrivateKey(kA, in, inLen); + } + else { + ret = wc_MlKemKey_DecodePublicKey(kA, in, inLen); + } + mlkem_oom_armed = 0; + (void)wolfSSL_SetAllocators(mlkem_oom_mf, mlkem_oom_ff, mlkem_oom_rf); + return ret; +} +#endif + +/* A decode into a full key that fails to allocate leaves the key unusable, + * not flagged as set over a freed buffer. */ +int test_wc_mlkem_decode_alloc_fail(void) +{ + EXPECT_DECLS; +#ifdef MLKEM_OOM_TEST + MlKemKey* kA = NULL; + MlKemKey* kB = NULL; + WC_RNG rng; + byte pkB[WC_ML_KEM_MAX_PUBLIC_KEY_SIZE]; + byte skB[WC_ML_KEM_MAX_PRIVATE_KEY_SIZE]; + byte ct[WC_ML_KEM_MAX_CIPHER_TEXT_SIZE]; + byte ss[WC_ML_KEM_SS_SZ]; + word32 pubLen = 0; + word32 privLen = 0; + word32 ctLen = 0; + int priv; +#ifndef WOLFSSL_NO_ML_KEM_768 + const int mlkemType = WC_ML_KEM_768; + const int k = WC_ML_KEM_768_K; +#elif !defined(WOLFSSL_NO_ML_KEM_512) + const int mlkemType = WC_ML_KEM_512; + const int k = WC_ML_KEM_512_K; +#else + const int mlkemType = WC_ML_KEM_1024; + const int k = WC_ML_KEM_1024_K; +#endif + + XMEMSET(&rng, 0, sizeof(rng)); + mlkem_oom_size = (size_t)k * MLKEM_N * sizeof(sword16); + ExpectIntEQ(wc_InitRng(&rng), 0); + ExpectNotNull(kA = (MlKemKey*)XMALLOC(sizeof(*kA), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (kA != NULL) { + XMEMSET(kA, 0, sizeof(*kA)); + } + ExpectNotNull(kB = (MlKemKey*)XMALLOC(sizeof(*kB), NULL, + DYNAMIC_TYPE_TMP_BUFFER)); + if (kB != NULL) { + XMEMSET(kB, 0, sizeof(*kB)); + } + ExpectIntEQ(wc_MlKemKey_Init(kA, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_Init(kB, mlkemType, NULL, INVALID_DEVID), 0); + ExpectIntEQ(wc_MlKemKey_MakeKey(kB, &rng), 0); + ExpectIntEQ(wc_MlKemKey_PublicKeySize(kB, &pubLen), 0); + ExpectIntEQ(wc_MlKemKey_PrivateKeySize(kB, &privLen), 0); + ExpectIntEQ(wc_MlKemKey_CipherTextSize(kB, &ctLen), 0); + ExpectIntEQ(wc_MlKemKey_EncodePublicKey(kB, pkB, pubLen), 0); + ExpectIntEQ(wc_MlKemKey_EncodePrivateKey(kB, skB, privLen), 0); + ExpectIntEQ(wc_MlKemKey_Encapsulate(kB, ct, ss, &rng), 0); + + for (priv = 0; priv <= 1; priv++) { + ExpectIntEQ(wc_MlKemKey_MakeKey(kA, &rng), 0); + if (EXPECT_SUCCESS()) { + ExpectIntEQ(mlkem_oom_decode(kA, priv, priv ? skB : pkB, + priv ? privLen : pubLen), WC_NO_ERR_TRACE(MEMORY_E)); + } + ExpectIntEQ(wc_MlKemKey_Encapsulate(kA, ct, ss, &rng), + WC_NO_ERR_TRACE(BAD_STATE_E)); + ExpectIntEQ(wc_MlKemKey_Decapsulate(kA, ss, ct, ctLen), + WC_NO_ERR_TRACE(BAD_STATE_E)); + } + + wc_MlKemKey_Free(kB); + wc_MlKemKey_Free(kA); + XFREE(kB, NULL, DYNAMIC_TYPE_TMP_BUFFER); + XFREE(kA, NULL, DYNAMIC_TYPE_TMP_BUFFER); + DoExpectIntEQ(wc_FreeRng(&rng), 0); +#endif + return EXPECT_RESULT(); +} diff --git a/tests/api/test_mlkem.h b/tests/api/test_mlkem.h index 1e89e89c31f..07d6e44c155 100644 --- a/tests/api/test_mlkem.h +++ b/tests/api/test_mlkem.h @@ -39,6 +39,9 @@ int test_wc_mlkem_encapsulate_pubkey_unset_decision(void); int test_wc_mlkem_encode_key_len_decision(void); int test_wc_mlkem_cb_free(void); int test_wc_mlkem_cb_pending_rejected(void); +int test_wc_mlkem_decode_pubkey_refused_save(void); +int test_wc_mlkem_decode_reused_key(void); +int test_wc_mlkem_decode_alloc_fail(void); #define TEST_MLKEM_DECLS \ TEST_DECL_GROUP("mlkem", test_wc_mlkem_make_key_kats), \ @@ -55,6 +58,9 @@ int test_wc_mlkem_cb_pending_rejected(void); TEST_DECL_GROUP("mlkem", test_wc_mlkem_encapsulate_pubkey_unset_decision), \ TEST_DECL_GROUP("mlkem", test_wc_mlkem_encode_key_len_decision), \ TEST_DECL_GROUP("mlkem", test_wc_mlkem_cb_free), \ - TEST_DECL_GROUP("mlkem", test_wc_mlkem_cb_pending_rejected) + TEST_DECL_GROUP("mlkem", test_wc_mlkem_cb_pending_rejected), \ + TEST_DECL_GROUP("mlkem", test_wc_mlkem_decode_pubkey_refused_save), \ + TEST_DECL_GROUP("mlkem", test_wc_mlkem_decode_reused_key), \ + TEST_DECL_GROUP("mlkem", test_wc_mlkem_decode_alloc_fail) #endif /* WOLFCRYPT_TEST_MLKEM_H */ diff --git a/tests/unit-mcdc/test_sp_x86_64_whitebox.c b/tests/unit-mcdc/test_sp_x86_64_whitebox.c index 2f561414ee2..0594ea1b4e0 100644 --- a/tests/unit-mcdc/test_sp_x86_64_whitebox.c +++ b/tests/unit-mcdc/test_sp_x86_64_whitebox.c @@ -151,25 +151,15 @@ * from this file's fixed/small-scalar inputs. */ -/* The richest dispatches here are four operands: - * - * IS_INTEL_BMI2(f) && IS_INTEL_ADX(f) && IS_INTEL_AVX2(f) && - * (SAVE_VECTOR_REGISTERS2() == 0) - * - * The feature bits are handled by the one-at-a-time masks in main(), but the - * save operand cannot be flipped that way: in a userspace build types.h - * resolves SAVE_VECTOR_REGISTERS2() to the literal 0, so "(0 == 0)" is - * structurally true and has no false side at all. It is real where the save - * can be refused (the kernel-module build). WC_CHECK_FOR_INTR_SIGNALS is the - * #ifndef extension point types.h offers for that, so defining it here -- - * before the .c below pulls in any wolfSSL header -- routes every - * SAVE_VECTOR_REGISTERS2() site through a variable this file controls. Same - * arrangement as test_wc_mlkem_poly_whitebox.c. */ /* Sweep depth for the allocation-failure pass. Each index repeats the * whole dispatch+crafted driving, and TEST_TIMEOUT is wall clock under * MAXPAR, so this stays modest. */ #define WB_FAULT_MAX_N 20 +/* CPUID alone picks each lane; the lane then takes the vector-register save + * and a refused save is an error, never a switch of lane. Userspace resolves + * SAVE_VECTOR_REGISTERS2() to 0, so WC_CHECK_FOR_INTR_SIGNALS is defined here, + * before any wolfSSL header, to let this file refuse the save on demand. */ static int wb_intr_ret = 0; #define WC_CHECK_FOR_INTR_SIGNALS() (wb_intr_ret) @@ -199,6 +189,20 @@ static int wb_intr_ret = 0; #include static int wb_fail = 0; +/* Set while the save is refused. Every driven operation must fail then, so + * these two counters turn that pass from a coverage sweep into a check. */ +static int wb_expect_refusal = 0; +static int wb_contract_fail = 0; /* a refused save failed to stop a call */ +static long wb_observed = 0; /* instrumented calls compiled in here */ +static long wb_refused_ok = 0; /* failed as required */ +static long wb_refused_bad = 0; /* succeeded despite a refused save */ +#define WB_OUTCOME(ret) do { \ + wb_observed++; \ + if (wb_expect_refusal) { \ + if ((ret) == 0) wb_refused_bad++; else wb_refused_ok++; \ + } } while (0) +/* Evaluates the call once, records its outcome, yields it to the caller. */ +#define WB_CHECK(call) __extension__ ({ int wb_r_ = (call); WB_OUTCOME(wb_r_); wb_r_; }) #define WB_NOTE(msg) do { printf(" [wb] %s\n", (msg)); } while (0) /* Crafted-input driver shared with the sp_c64.c/sp_c32.c white-boxes: the @@ -270,12 +274,12 @@ static void wb_run_ecc_curve(int curve_id, int fieldSz, const char* label) return; } - if (wc_ecc_make_key_ex(&rng, fieldSz, &keyA, curve_id) != 0) { + if (WB_CHECK(wc_ecc_make_key_ex(&rng, fieldSz, &keyA, curve_id)) != 0) { WB_NOTE("wc_ecc_make_key_ex(keyA) failed"); wb_fail = 1; ok = 0; } - if (ok && wc_ecc_make_key_ex(&rng, fieldSz, &keyB, curve_id) != 0) { + if (ok && WB_CHECK(wc_ecc_make_key_ex(&rng, fieldSz, &keyB, curve_id)) != 0) { WB_NOTE("wc_ecc_make_key_ex(keyB) failed"); wb_fail = 1; ok = 0; @@ -283,20 +287,20 @@ static void wb_run_ecc_curve(int curve_id, int fieldSz, const char* label) if (ok) { sigLen = (word32)sizeof(sig); - if (wc_ecc_sign_hash(wb_digest, (word32)sizeof(wb_digest), sig, - &sigLen, &rng, &keyA) != 0) { + if (WB_CHECK(wc_ecc_sign_hash(wb_digest, (word32)sizeof(wb_digest), sig, + &sigLen, &rng, &keyA)) != 0) { WB_NOTE("wc_ecc_sign_hash failed"); wb_fail = 1; } - else if (wc_ecc_verify_hash(sig, sigLen, wb_digest, - (word32)sizeof(wb_digest), &verifyRes, &keyA) != 0) { + else if (WB_CHECK(wc_ecc_verify_hash(sig, sigLen, wb_digest, + (word32)sizeof(wb_digest), &verifyRes, &keyA)) != 0) { WB_NOTE("wc_ecc_verify_hash failed"); wb_fail = 1; } PRIVATE_KEY_UNLOCK(); secretALen = (word32)sizeof(secretA); - if (wc_ecc_shared_secret(&keyA, &keyB, secretA, &secretALen) != 0) { + if (WB_CHECK(wc_ecc_shared_secret(&keyA, &keyB, secretA, &secretALen)) != 0) { WB_NOTE("wc_ecc_shared_secret(A,B) failed"); wb_fail = 1; } @@ -321,6 +325,93 @@ static void wb_run_ecc_curve(int curve_id, int fieldSz, const char* label) #endif } +/* A refused save must stop work that uses vector registers and must not stop + * work that does not. On the base lane the only xmm user is the + * cache-resistant table lookup, which runs when ct is set, and ECDSA verify + * passes ct == 0. Caller clears AVX2 so the base lane is the one driven. */ +static void wb_run_ecc_verify_no_save(int curve_id, int fieldSz, + const char* label) +{ +#if defined(HAVE_ECC_SIGN) && defined(HAVE_ECC_VERIFY) + ecc_key keyA; + WC_RNG rng; + byte sig[ECC_MAX_SIG_SIZE]; + word32 sigLen = (word32)sizeof(sig); + int verifyRes = 0; + int ret; + + XMEMSET(&keyA, 0, sizeof(keyA)); + XMEMSET(&rng, 0, sizeof(rng)); + XMEMSET(sig, 0, sizeof(sig)); + + if (wc_ecc_init(&keyA) != 0) { + WB_NOTE("wc_ecc_init failed (verify without a save)"); + wb_fail = 1; + return; + } + if (wc_InitRng(&rng) != 0) { + WB_NOTE("wc_InitRng failed (verify without a save)"); + wb_fail = 1; + wc_ecc_free(&keyA); + return; + } + + /* Key and signature are made with the save allowed. */ + if (wc_ecc_make_key_ex(&rng, fieldSz, &keyA, curve_id) != 0) { + WB_NOTE("wc_ecc_make_key_ex failed (verify without a save)"); + wb_fail = 1; + } + else if (wc_ecc_sign_hash(wb_digest, (word32)sizeof(wb_digest), sig, + &sigLen, &rng, &keyA) != 0) { + WB_NOTE("wc_ecc_sign_hash failed (verify without a save)"); + wb_fail = 1; + } + else { + wb_intr_ret = 1; + ret = wc_ecc_verify_hash(sig, sigLen, wb_digest, + (word32)sizeof(wb_digest), &verifyRes, &keyA); + wb_intr_ret = 0; + + if (ret != 0) { + printf(" [wb] FAIL: %s stopped on a save it never needed\n", + label); + wb_contract_fail = 1; + } + else if (verifyRes != 1) { + printf(" [wb] FAIL: %s rejected a good signature\n", label); + wb_contract_fail = 1; + } + else { + WB_NOTE(label); + } + } + + wc_FreeRng(&rng); + wc_ecc_free(&keyA); +#else + (void)curve_id; + (void)fieldSz; + (void)label; + WB_NOTE("HAVE_ECC_SIGN/VERIFY not both defined; verify-without-save skipped"); +#endif +} + +static void wb_run_ecc_no_save(void) +{ +#ifndef WOLFSSL_SP_NO_256 + wb_run_ecc_verify_no_save(ECC_SECP256R1, 32, + "P-256 verify with the save refused on the base lane"); +#endif +#ifdef WOLFSSL_SP_384 + wb_run_ecc_verify_no_save(ECC_SECP384R1, 48, + "P-384 verify with the save refused on the base lane"); +#endif +#ifdef WOLFSSL_SP_521 + wb_run_ecc_verify_no_save(ECC_SECP521R1, 66, + "P-521 verify with the save refused on the base lane"); +#endif +} + static void wb_run_ecc(void) { #ifndef WOLFSSL_SP_NO_256 @@ -349,6 +440,10 @@ static void wb_run_ecc(void) { WB_NOTE("WOLFSSL_HAVE_SP_ECC/HAVE_ECC not both defined; ECC skipped"); } +static void wb_run_ecc_no_save(void) +{ + WB_NOTE("WOLFSSL_HAVE_SP_ECC/HAVE_ECC not both defined; ECC skipped"); +} #endif /* WOLFSSL_HAVE_SP_ECC && HAVE_ECC */ #if defined(WOLFSSL_HAVE_SP_RSA) && !defined(NO_RSA) && \ @@ -1154,7 +1249,7 @@ static void wb_run_dispatch_256(void) XMEMSET(tmp2, 0, sizeof(tmp2)); pp1.x[0] = 1; pp1.y[0] = 1; pp1.z[0] = 1; pp2.x[0] = 1; pp2.y[0] = 1; pp2.z[0] = 1; - sp_256_add_points_4(&pp1, &pp2, tmp2); + (void)WB_CHECK(sp_256_add_points_4(&pp1, &pp2, tmp2)); } { sp_point_256 pt; @@ -1291,7 +1386,7 @@ static void wb_run_dispatch_384(void) XMEMSET(tmp2, 0, sizeof(tmp2)); pp1.x[0] = 1; pp1.y[0] = 1; pp1.z[0] = 1; pp2.x[0] = 1; pp2.y[0] = 1; pp2.z[0] = 1; - sp_384_add_points_6(&pp1, &pp2, tmp2); + (void)WB_CHECK(sp_384_add_points_6(&pp1, &pp2, tmp2)); } { sp_point_384 pt; @@ -1428,7 +1523,7 @@ static void wb_run_dispatch_521(void) XMEMSET(tmp2, 0, sizeof(tmp2)); pp1.x[0] = 1; pp1.y[0] = 1; pp1.z[0] = 1; pp2.x[0] = 1; pp2.y[0] = 1; pp2.z[0] = 1; - sp_521_add_points_9(&pp1, &pp2, tmp2); + (void)WB_CHECK(sp_521_add_points_9(&pp1, &pp2, tmp2)); } { sp_point_521 pt; @@ -1589,12 +1684,12 @@ static void wb_run_crafted_curve(int curve_id, int fieldSz, return; } - if (wc_ecc_make_key_ex(&rng, fieldSz, &keyA, curve_id) != 0) { + if (WB_CHECK(wc_ecc_make_key_ex(&rng, fieldSz, &keyA, curve_id)) != 0) { WB_NOTE("wc_ecc_make_key_ex(keyA) failed (crafted)"); wb_fail = 1; ok = 0; } - if (ok && wc_ecc_make_key_ex(&rng, fieldSz, &keyB, curve_id) != 0) { + if (ok && WB_CHECK(wc_ecc_make_key_ex(&rng, fieldSz, &keyB, curve_id)) != 0) { WB_NOTE("wc_ecc_make_key_ex(keyB) failed (crafted)"); wb_fail = 1; ok = 0; @@ -1995,18 +2090,46 @@ int main(void) wb_run_crafted(); wb_spc_all(); - /* Fourth operand: every feature present but the vector-register save - * refused, so each chain falls through on its last condition. */ + /* Refused save: every lane returns its error instead of running, so + * the drivers report failures here by design. The counters turn that + * into a check: a call that SUCCEEDS with the save refused means the + * dispatch found another way to run, which is what this file exists + * to keep out. */ cpuid_select_flags(real); wb_intr_ret = 1; + wb_expect_refusal = 1; wb_run_ecc(); wb_run_rsa_signverify(); wb_run_dh(); wb_run_dispatch(); wb_run_crafted(); wb_spc_all(); + wb_expect_refusal = 0; wb_intr_ret = 0; + printf(" [wb] refused save: %ld calls failed as required, %ld ran anyway\n", + wb_refused_ok, wb_refused_bad); + if (wb_refused_bad != 0) { + printf(" [wb] FAIL: a refused vector-register save did not stop the call\n"); + wb_contract_fail = 1; + } + /* The refusal pass compiles for SP RSA or DH too, but the instrumented + * calls are all ECC. In a build without them there is nothing to + * check, which is not the same as a check that failed. */ + if (wb_observed == 0) { + printf(" [wb] no instrumented call in this configuration; nothing to check\n"); + } + else if (wb_refused_ok == 0) { + printf(" [wb] FAIL: nothing was seen failing, so this check proves nothing\n"); + wb_contract_fail = 1; + } + + /* The other half of the contract: with AVX2 cleared, ECDSA verify + * needs no vector registers, so a refused save must not stop it. */ + cpuid_select_flags(real & ~(cpuid_flags_t)CPUID_AVX2); + wb_run_ecc_no_save(); + cpuid_select_flags(real); + wb_run_rsa_free(); /* Allocation-failure pass. @@ -2052,5 +2175,6 @@ int main(void) printf(" no SP feature; nothing to exercise\n"); #endif (void)wb_fail; - return 0; + /* Coverage sweeps stay advisory; the fail-closed contract does not. */ + return wb_contract_fail; } diff --git a/tests/unit-mcdc/test_wc_mldsa_whitebox.c b/tests/unit-mcdc/test_wc_mldsa_whitebox.c index 41409587cce..6d76d9f6e89 100644 --- a/tests/unit-mcdc/test_wc_mldsa_whitebox.c +++ b/tests/unit-mcdc/test_wc_mldsa_whitebox.c @@ -47,9 +47,9 @@ * the binary always returns 0 so the harness keeps the variant. */ -/* SAVE_VECTOR_REGISTERS2() gates every SIMD dispatch in this file. In a - * userspace build types.h resolves it to the literal 0, so "(0 == 0)" is - * structurally true and that operand has no false side at all -- it is real +/* SAVE_VECTOR_REGISTERS2() is taken by every SIMD lane in this file, and a + * refusal is returned as an error. In a userspace build types.h resolves it + * to the literal 0, so the refusal branch is unreachable -- it is real * only where the save can be refused (the kernel-module build, where it * becomes WC_CHECK_FOR_INTR_SIGNALS()). That is the #ifndef extension point * types.h offers, so defining it here -- BEFORE any wolfSSL header is reached @@ -331,7 +331,8 @@ static void wb_check_hint_inner_loops(void) /* ------------------------------------------------------------------ * * mldsa_make_hint_88 / _32 / mldsa_make_hint: the 3-way compound * (s>LOW) || (s<-LOW) || ((s==-LOW) && (w1!=0)) - * and the too-many-hints guard (idx>OMEGA -> -1), plus mldsa_make_hint's + * and the too-many-hints guard (idx>OMEGA -> *valid = 0), plus + * mldsa_make_hint's * gamma2 dispatch (88 arm / 32 arm / neither). * ------------------------------------------------------------------ */ #ifndef WOLFSSL_MLDSA_NO_SIGN @@ -342,6 +343,7 @@ static void wb_make_hint_88(void) sword32 w1[MLDSA_N]; byte h[256]; byte idx; + int valid; unsigned int j; int ret; const sword32 low = (sword32)MLDSA_Q_LOW_88; @@ -353,7 +355,7 @@ static void wb_make_hint_88(void) /* All three operands FALSE for every coefficient -> no hint, idx stays 0. */ idx = 0; - ret = mldsa_make_hint_88(s, w1, h, &idx); + ret = mldsa_make_hint_88(s, w1, h, &idx, &valid); if ((ret != 0) || (idx != 0)) { WB_NOTE("mldsa_make_hint_88(no-hint) unexpected"); } @@ -361,7 +363,7 @@ static void wb_make_hint_88(void) /* First operand TRUE (s > LOW). */ idx = 0; s[1] = low + 1; - ret = mldsa_make_hint_88(s, w1, h, &idx); + ret = mldsa_make_hint_88(s, w1, h, &idx, &valid); if ((ret != 0) || (idx != 1)) { WB_NOTE("mldsa_make_hint_88(s>LOW) unexpected"); } @@ -370,7 +372,7 @@ static void wb_make_hint_88(void) idx = 0; s[1] = 0; s[2] = -low - 1; - ret = mldsa_make_hint_88(s, w1, h, &idx); + ret = mldsa_make_hint_88(s, w1, h, &idx, &valid); if ((ret != 0) || (idx != 1)) { WB_NOTE("mldsa_make_hint_88(s<-LOW) unexpected"); } @@ -380,7 +382,7 @@ static void wb_make_hint_88(void) s[2] = 0; s[3] = -low; w1[3] = 1; - ret = mldsa_make_hint_88(s, w1, h, &idx); + ret = mldsa_make_hint_88(s, w1, h, &idx, &valid); if ((ret != 0) || (idx != 1)) { WB_NOTE("mldsa_make_hint_88(s==-LOW,w1!=0) unexpected"); } @@ -389,7 +391,7 @@ static void wb_make_hint_88(void) * (independence of the w1!=0 operand). */ idx = 0; w1[3] = 0; - ret = mldsa_make_hint_88(s, w1, h, &idx); + ret = mldsa_make_hint_88(s, w1, h, &idx, &valid); if ((ret != 0) || (idx != 0)) { WB_NOTE("mldsa_make_hint_88(s==-LOW,w1==0) unexpected"); } @@ -400,9 +402,9 @@ static void wb_make_hint_88(void) s[j] = low + 1; w1[j] = 0; } - ret = mldsa_make_hint_88(s, w1, h, &idx); - if (ret != -1) { - WB_NOTE("mldsa_make_hint_88(too-many) expected -1"); + ret = mldsa_make_hint_88(s, w1, h, &idx, &valid); + if ((ret != 0) || (valid != 0)) { + WB_NOTE("mldsa_make_hint_88(too-many) expected valid = 0"); } WB_OK("mldsa_make_hint_88 operand + overflow pairs exercised"); } @@ -415,6 +417,7 @@ static void wb_make_hint_32(void) sword32 w1[MLDSA_N]; byte h[256]; byte idx; + int valid; unsigned int j; int ret; const sword32 low = (sword32)MLDSA_Q_LOW_32; @@ -427,7 +430,7 @@ static void wb_make_hint_32(void) /* No hint. */ idx = 0; - ret = mldsa_make_hint_32(s, w1, omega, h, &idx); + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, &valid); if ((ret != 0) || (idx != 0)) { WB_NOTE("mldsa_make_hint_32(no-hint) unexpected"); } @@ -435,7 +438,7 @@ static void wb_make_hint_32(void) /* s > LOW. */ idx = 0; s[1] = low + 1; - ret = mldsa_make_hint_32(s, w1, omega, h, &idx); + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, &valid); if ((ret != 0) || (idx != 1)) { WB_NOTE("mldsa_make_hint_32(s>LOW) unexpected"); } @@ -444,7 +447,7 @@ static void wb_make_hint_32(void) idx = 0; s[1] = 0; s[2] = -low - 1; - ret = mldsa_make_hint_32(s, w1, omega, h, &idx); + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, &valid); if ((ret != 0) || (idx != 1)) { WB_NOTE("mldsa_make_hint_32(s<-LOW) unexpected"); } @@ -454,7 +457,7 @@ static void wb_make_hint_32(void) s[2] = 0; s[3] = -low; w1[3] = 1; - ret = mldsa_make_hint_32(s, w1, omega, h, &idx); + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, &valid); if ((ret != 0) || (idx != 1)) { WB_NOTE("mldsa_make_hint_32(s==-LOW,w1!=0) unexpected"); } @@ -462,7 +465,7 @@ static void wb_make_hint_32(void) /* (s == -LOW) && (w1 == 0) -> no hint. */ idx = 0; w1[3] = 0; - ret = mldsa_make_hint_32(s, w1, omega, h, &idx); + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, &valid); if ((ret != 0) || (idx != 0)) { WB_NOTE("mldsa_make_hint_32(s==-LOW,w1==0) unexpected"); } @@ -472,9 +475,9 @@ static void wb_make_hint_32(void) for (j = 0; j < MLDSA_N; j++) { s[j] = low + 1; } - ret = mldsa_make_hint_32(s, w1, omega, h, &idx); - if (ret != -1) { - WB_NOTE("mldsa_make_hint_32(too-many) expected -1"); + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, &valid); + if ((ret != 0) || (valid != 0)) { + WB_NOTE("mldsa_make_hint_32(too-many) expected valid = 0"); } WB_OK("mldsa_make_hint_32 operand + overflow pairs exercised"); } @@ -939,7 +942,7 @@ static void wb_dispatch_rows(void) { CPUID_INTEL, CPUID_AVX512_BW, 0 }, /* F set, BW clear */ { CPUID_INTEL, CPUID_AVX512_BW, 1 }, /* BW clear, refuse */ /* The AVX512 matrix/mask dispatches read - * USE_INTEL_AVX512(f) && IS_INTEL_BMI2(f) && (save == 0) + * USE_INTEL_AVX512(f) && IS_INTEL_BMI2(f) * so their BMI2 operand only takes its false side on a row that keeps * AVX512 and drops BMI2 -- dropping both together (further down) * never evaluates it. */ diff --git a/tests/unit-mcdc/test_wc_mlkem_poly_whitebox.c b/tests/unit-mcdc/test_wc_mlkem_poly_whitebox.c index 28bdd087790..469e2d911de 100644 --- a/tests/unit-mcdc/test_wc_mlkem_poly_whitebox.c +++ b/tests/unit-mcdc/test_wc_mlkem_poly_whitebox.c @@ -60,9 +60,8 @@ * exactly that, so defining it here -- BEFORE any wolfSSL header is pulled in * by the .c below -- routes all 58 SAVE_VECTOR_REGISTERS2() sites through a * variable this file controls, using the library's own hook rather than - * overriding a macro behind its back. Setting it non-zero makes each dispatch - * fall through to the portable C path, which is what the operand's false side - * selects on a platform that really can refuse. + * overriding a macro behind its back. Setting it non-zero refuses the save, + * which each dispatch reports as an error rather than switching lanes. * * The file uses only SAVE_VECTOR_REGISTERS2(); the SAVE_VECTOR_REGISTERS( * fail_clause) form, whose expansion also changes under this hook, appears @@ -72,14 +71,14 @@ * the save-refused answer at a CHOSEN call index instead of only "always" or * "never". Two dispatch guards nest inside one another: * - * mlkem_gen_matrix() if (IS_INTEL_AVX2(..) && save == 0) + * mlkem_gen_matrix() if (IS_INTEL_AVX2(..)) save * mlkem_gen_matrix_k3_avx2() for (..) { if (IS_INTEL_BMI2(..)) - * else if (IS_INTEL_AVX2(..) - * && save == 0) + * else if (IS_INTEL_AVX2(..)) + * save * - * so the inner guard's operands can only be reached when the OUTER one was - * already satisfied. A process-wide "always refuse" therefore never lets the - * inner site run at all, and its false side would stay unreachable. Letting + * so the inner save can only be reached when the OUTER one was already + * granted. A process-wide "always refuse" therefore never lets the inner + * site run at all, and its refusal branch would stay unreachable. Letting * call 0..n-1 succeed and calls >= n refuse gives the inner site a genuine * (T,F) row while the outer one still took (T,T). * diff --git a/wolfcrypt/src/cpuid.c b/wolfcrypt/src/cpuid.c index 39b45818b36..9d9e05bdd1f 100644 --- a/wolfcrypt/src/cpuid.c +++ b/wolfcrypt/src/cpuid.c @@ -992,6 +992,28 @@ return WOLFSSL_ATOMIC_LOAD(cpuid_flags); } + /* FIPS v7 reads the CPU features once at power on and ignores changes. + * Linux only warns when microcode changes features while running and + * says to reboot (Documentation/arch/x86/microcode.rst, late loading); + * Intel marks such updates unfit to load while running ("Minimum Runtime + * Microcode Update Revision"). A reboot re-reads the features. */ +#if defined(HAVE_FIPS) && FIPS_VERSION3_GE(7,0,0) && \ + !defined(WOLFSSL_FIPS_DEV) && !defined(WOLFSSL_FIPS_DEV_NO_POST) + void cpuid_select_flags(cpuid_flags_t flags) + { + (void)flags; + } + + void cpuid_set_flag(cpuid_flags_t flag) + { + (void)flag; + } + + void cpuid_clear_flag(cpuid_flags_t flag) + { + (void)flag; + } +#else void cpuid_select_flags(cpuid_flags_t flags) { WOLFSSL_ATOMIC_STORE(cpuid_flags, flags); @@ -1012,5 +1034,6 @@ (&cpuid_flags, ¤t_flags, current_flags & ~flag)) WC_RELAX_LONG_LOOP(); } +#endif #endif /* HAVE_CPUID */ diff --git a/wolfcrypt/src/sp_x86_64.c b/wolfcrypt/src/sp_x86_64.c index c5bb7fda462..e82c59681ce 100644 --- a/wolfcrypt/src/sp_x86_64.c +++ b/wolfcrypt/src/sp_x86_64.c @@ -2037,6 +2037,7 @@ int sp_RsaPublic_2048(const byte* in, word32 inLen, const mp_int* em, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif if (*outLen < 256) { @@ -2071,10 +2072,14 @@ int sp_RsaPublic_2048(const byte* in, word32 inLen, const mp_int* em, sp_2048_from_mp(m, 32, mm); #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } #endif if (e == 0x10001) { @@ -2085,13 +2090,15 @@ int sp_RsaPublic_2048(const byte* in, word32 inLen, const mp_int* em, /* Convert to Montgomery form. */ XMEMSET(a, 0, sizeof(sp_digit) * 32); - err = sp_2048_mod_32_cond(r, a, m); + /* Keep a failed save's error. */ + if (err == MP_OKAY) + err = sp_2048_mod_32_cond(r, a, m); /* Montgomery form: r = a.R mod m */ if (err == MP_OKAY) { /* r = a ^ 0x10000 => r = a squared 16 times */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { for (i = 15; i >= 0; i--) { sp_2048_mont_sqr_avx2_32(r, r, m, mp); } @@ -2122,7 +2129,7 @@ int sp_RsaPublic_2048(const byte* in, word32 inLen, const mp_int* em, } else if (e == 0x3) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { if (err == MP_OKAY) { sp_2048_sqr_avx2_32(r, ah); err = sp_2048_mod_32_cond(r, r, m); @@ -2153,7 +2160,9 @@ int sp_RsaPublic_2048(const byte* in, word32 inLen, const mp_int* em, /* Convert to Montgomery form. */ XMEMSET(a, 0, sizeof(sp_digit) * 32); - err = sp_2048_mod_32_cond(a, a, m); + /* Keep a failed save's error. */ + if (err == MP_OKAY) + err = sp_2048_mod_32_cond(a, a, m); if (err == MP_OKAY) { for (i=63; i>=0; i--) { @@ -2164,7 +2173,7 @@ int sp_RsaPublic_2048(const byte* in, word32 inLen, const mp_int* em, XMEMCPY(r, a, sizeof(sp_digit) * 32); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { for (i--; i>=0; i--) { sp_2048_mont_sqr_avx2_32(r, r, m, mp); if (((e >> i) & 1) == 1) { @@ -2347,6 +2356,7 @@ int sp_RsaPrivate_2048(const byte* in, word32 inLen, const mp_int* dm, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif (void)dm; @@ -2384,22 +2394,32 @@ int sp_RsaPrivate_2048(const byte* in, word32 inLen, const mp_int* dm, sp_2048_from_mp(dp, 16, dpm); #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif - if (saved_vector_registers) - err = sp_2048_mod_exp_avx2_16(tmpa, a, dp, 1024, p, 1); - else + /* err may hold a failed save. */ + if (err == MP_OKAY) { +#ifdef HAVE_INTEL_AVX2 + if (use_avx2_lane) + err = sp_2048_mod_exp_avx2_16(tmpa, a, dp, 1024, p, 1); + else #endif - err = sp_2048_mod_exp_16(tmpa, a, dp, 1024, p, 1); + err = sp_2048_mod_exp_16(tmpa, a, dp, 1024, p, 1); + } } if (err == MP_OKAY) { sp_2048_from_mp(dq, 16, dqm); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) err = sp_2048_mod_exp_avx2_16(tmpb, a, dq, 1024, q, 1); - else + else #endif err = sp_2048_mod_exp_16(tmpb, a, dq, 1024, q, 1); } @@ -2407,7 +2427,7 @@ int sp_RsaPrivate_2048(const byte* in, word32 inLen, const mp_int* dm, if (err == MP_OKAY) { c = sp_2048_sub_in_place_16(tmpa, tmpb); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { c += sp_2048_cond_add_avx2_16(tmpa, tmpa, p, c); sp_2048_cond_add_avx2_16(tmpa, tmpa, p, c); } @@ -2420,7 +2440,7 @@ int sp_RsaPrivate_2048(const byte* in, word32 inLen, const mp_int* dm, sp_2048_from_mp(qi, 16, qim); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_2048_mul_avx2_16(tmpa, tmpa, qi); else #endif @@ -2430,7 +2450,7 @@ int sp_RsaPrivate_2048(const byte* in, word32 inLen, const mp_int* dm, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_2048_mul_avx2_16(tmpa, q, tmpa); else #endif @@ -2573,9 +2593,12 @@ int sp_ModExp_2048(const mp_int* base, const mp_int* exp, const mp_int* mod, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_2048_mod_exp_avx2_32(r, b, e, expBits, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_2048_mod_exp_avx2_32(r, b, e, expBits, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -2897,10 +2920,12 @@ int sp_DhExp_2048(const mp_int* base, const byte* exp, word32 expLen, if (base->used == 1 && base->dp[0] == 2 && m[31] == (sp_digit)-1) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_2048_mod_exp_2_avx2_32(r, e, (int)expLen * 8, m); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_2048_mod_exp_2_avx2_32(r, e, (int)expLen * 8, m); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -2911,10 +2936,12 @@ int sp_DhExp_2048(const mp_int* base, const byte* exp, word32 expLen, { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_2048_mod_exp_avx2_32(r, b, e, (int)expLen * 8, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_2048_mod_exp_avx2_32(r, b, e, (int)expLen * 8, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -2988,9 +3015,12 @@ int sp_ModExp_1024(const mp_int* base, const mp_int* exp, const mp_int* mod, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_2048_mod_exp_avx2_16(r, b, e, expBits, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_2048_mod_exp_avx2_16(r, b, e, expBits, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -4821,6 +4851,7 @@ int sp_RsaPublic_3072(const byte* in, word32 inLen, const mp_int* em, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif if (*outLen < 384) { @@ -4855,10 +4886,14 @@ int sp_RsaPublic_3072(const byte* in, word32 inLen, const mp_int* em, sp_3072_from_mp(m, 48, mm); #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } #endif if (e == 0x10001) { @@ -4869,13 +4904,15 @@ int sp_RsaPublic_3072(const byte* in, word32 inLen, const mp_int* em, /* Convert to Montgomery form. */ XMEMSET(a, 0, sizeof(sp_digit) * 48); - err = sp_3072_mod_48_cond(r, a, m); + /* Keep a failed save's error. */ + if (err == MP_OKAY) + err = sp_3072_mod_48_cond(r, a, m); /* Montgomery form: r = a.R mod m */ if (err == MP_OKAY) { /* r = a ^ 0x10000 => r = a squared 16 times */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { for (i = 15; i >= 0; i--) { sp_3072_mont_sqr_avx2_48(r, r, m, mp); } @@ -4906,7 +4943,7 @@ int sp_RsaPublic_3072(const byte* in, word32 inLen, const mp_int* em, } else if (e == 0x3) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { if (err == MP_OKAY) { sp_3072_sqr_avx2_48(r, ah); err = sp_3072_mod_48_cond(r, r, m); @@ -4937,7 +4974,9 @@ int sp_RsaPublic_3072(const byte* in, word32 inLen, const mp_int* em, /* Convert to Montgomery form. */ XMEMSET(a, 0, sizeof(sp_digit) * 48); - err = sp_3072_mod_48_cond(a, a, m); + /* Keep a failed save's error. */ + if (err == MP_OKAY) + err = sp_3072_mod_48_cond(a, a, m); if (err == MP_OKAY) { for (i=63; i>=0; i--) { @@ -4948,7 +4987,7 @@ int sp_RsaPublic_3072(const byte* in, word32 inLen, const mp_int* em, XMEMCPY(r, a, sizeof(sp_digit) * 48); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { for (i--; i>=0; i--) { sp_3072_mont_sqr_avx2_48(r, r, m, mp); if (((e >> i) & 1) == 1) { @@ -5131,6 +5170,7 @@ int sp_RsaPrivate_3072(const byte* in, word32 inLen, const mp_int* dm, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif (void)dm; @@ -5168,22 +5208,32 @@ int sp_RsaPrivate_3072(const byte* in, word32 inLen, const mp_int* dm, sp_3072_from_mp(dp, 24, dpm); #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif - if (saved_vector_registers) - err = sp_3072_mod_exp_avx2_24(tmpa, a, dp, 1536, p, 1); - else + /* err may hold a failed save. */ + if (err == MP_OKAY) { +#ifdef HAVE_INTEL_AVX2 + if (use_avx2_lane) + err = sp_3072_mod_exp_avx2_24(tmpa, a, dp, 1536, p, 1); + else #endif - err = sp_3072_mod_exp_24(tmpa, a, dp, 1536, p, 1); + err = sp_3072_mod_exp_24(tmpa, a, dp, 1536, p, 1); + } } if (err == MP_OKAY) { sp_3072_from_mp(dq, 24, dqm); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) err = sp_3072_mod_exp_avx2_24(tmpb, a, dq, 1536, q, 1); - else + else #endif err = sp_3072_mod_exp_24(tmpb, a, dq, 1536, q, 1); } @@ -5191,7 +5241,7 @@ int sp_RsaPrivate_3072(const byte* in, word32 inLen, const mp_int* dm, if (err == MP_OKAY) { c = sp_3072_sub_in_place_24(tmpa, tmpb); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { c += sp_3072_cond_add_avx2_24(tmpa, tmpa, p, c); sp_3072_cond_add_avx2_24(tmpa, tmpa, p, c); } @@ -5204,7 +5254,7 @@ int sp_RsaPrivate_3072(const byte* in, word32 inLen, const mp_int* dm, sp_3072_from_mp(qi, 24, qim); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_3072_mul_avx2_24(tmpa, tmpa, qi); else #endif @@ -5214,7 +5264,7 @@ int sp_RsaPrivate_3072(const byte* in, word32 inLen, const mp_int* dm, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_3072_mul_avx2_24(tmpa, q, tmpa); else #endif @@ -5357,9 +5407,12 @@ int sp_ModExp_3072(const mp_int* base, const mp_int* exp, const mp_int* mod, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_3072_mod_exp_avx2_48(r, b, e, expBits, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_3072_mod_exp_avx2_48(r, b, e, expBits, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -5681,10 +5734,12 @@ int sp_DhExp_3072(const mp_int* base, const byte* exp, word32 expLen, if (base->used == 1 && base->dp[0] == 2 && m[47] == (sp_digit)-1) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_3072_mod_exp_2_avx2_48(r, e, (int)expLen * 8, m); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_3072_mod_exp_2_avx2_48(r, e, (int)expLen * 8, m); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -5695,10 +5750,12 @@ int sp_DhExp_3072(const mp_int* base, const byte* exp, word32 expLen, { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_3072_mod_exp_avx2_48(r, b, e, (int)expLen * 8, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_3072_mod_exp_avx2_48(r, b, e, (int)expLen * 8, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -5772,9 +5829,12 @@ int sp_ModExp_1536(const mp_int* base, const mp_int* exp, const mp_int* mod, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_3072_mod_exp_avx2_24(r, b, e, expBits, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_3072_mod_exp_avx2_24(r, b, e, expBits, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -6832,6 +6892,7 @@ int sp_RsaPublic_4096(const byte* in, word32 inLen, const mp_int* em, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif if (*outLen < 512) { @@ -6866,10 +6927,14 @@ int sp_RsaPublic_4096(const byte* in, word32 inLen, const mp_int* em, sp_4096_from_mp(m, 64, mm); #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } #endif if (e == 0x10001) { @@ -6880,13 +6945,15 @@ int sp_RsaPublic_4096(const byte* in, word32 inLen, const mp_int* em, /* Convert to Montgomery form. */ XMEMSET(a, 0, sizeof(sp_digit) * 64); - err = sp_4096_mod_64_cond(r, a, m); + /* Keep a failed save's error. */ + if (err == MP_OKAY) + err = sp_4096_mod_64_cond(r, a, m); /* Montgomery form: r = a.R mod m */ if (err == MP_OKAY) { /* r = a ^ 0x10000 => r = a squared 16 times */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { for (i = 15; i >= 0; i--) { sp_4096_mont_sqr_avx2_64(r, r, m, mp); } @@ -6917,7 +6984,7 @@ int sp_RsaPublic_4096(const byte* in, word32 inLen, const mp_int* em, } else if (e == 0x3) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { if (err == MP_OKAY) { sp_4096_sqr_avx2_64(r, ah); err = sp_4096_mod_64_cond(r, r, m); @@ -6948,7 +7015,9 @@ int sp_RsaPublic_4096(const byte* in, word32 inLen, const mp_int* em, /* Convert to Montgomery form. */ XMEMSET(a, 0, sizeof(sp_digit) * 64); - err = sp_4096_mod_64_cond(a, a, m); + /* Keep a failed save's error. */ + if (err == MP_OKAY) + err = sp_4096_mod_64_cond(a, a, m); if (err == MP_OKAY) { for (i=63; i>=0; i--) { @@ -6959,7 +7028,7 @@ int sp_RsaPublic_4096(const byte* in, word32 inLen, const mp_int* em, XMEMCPY(r, a, sizeof(sp_digit) * 64); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { for (i--; i>=0; i--) { sp_4096_mont_sqr_avx2_64(r, r, m, mp); if (((e >> i) & 1) == 1) { @@ -7142,6 +7211,7 @@ int sp_RsaPrivate_4096(const byte* in, word32 inLen, const mp_int* dm, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif (void)dm; @@ -7179,22 +7249,32 @@ int sp_RsaPrivate_4096(const byte* in, word32 inLen, const mp_int* dm, sp_4096_from_mp(dp, 32, dpm); #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif - if (saved_vector_registers) - err = sp_2048_mod_exp_avx2_32(tmpa, a, dp, 2048, p, 1); - else + /* err may hold a failed save. */ + if (err == MP_OKAY) { +#ifdef HAVE_INTEL_AVX2 + if (use_avx2_lane) + err = sp_2048_mod_exp_avx2_32(tmpa, a, dp, 2048, p, 1); + else #endif - err = sp_2048_mod_exp_32(tmpa, a, dp, 2048, p, 1); + err = sp_2048_mod_exp_32(tmpa, a, dp, 2048, p, 1); + } } if (err == MP_OKAY) { sp_4096_from_mp(dq, 32, dqm); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) err = sp_2048_mod_exp_avx2_32(tmpb, a, dq, 2048, q, 1); - else + else #endif err = sp_2048_mod_exp_32(tmpb, a, dq, 2048, q, 1); } @@ -7202,7 +7282,7 @@ int sp_RsaPrivate_4096(const byte* in, word32 inLen, const mp_int* dm, if (err == MP_OKAY) { c = sp_2048_sub_in_place_32(tmpa, tmpb); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (use_avx2_lane) { c += sp_4096_cond_add_avx2_32(tmpa, tmpa, p, c); sp_4096_cond_add_avx2_32(tmpa, tmpa, p, c); } @@ -7215,7 +7295,7 @@ int sp_RsaPrivate_4096(const byte* in, word32 inLen, const mp_int* dm, sp_2048_from_mp(qi, 32, qim); #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_2048_mul_avx2_32(tmpa, tmpa, qi); else #endif @@ -7225,7 +7305,7 @@ int sp_RsaPrivate_4096(const byte* in, word32 inLen, const mp_int* dm, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_2048_mul_avx2_32(tmpa, q, tmpa); else #endif @@ -7368,9 +7448,12 @@ int sp_ModExp_4096(const mp_int* base, const mp_int* exp, const mp_int* mod, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_4096_mod_exp_avx2_64(r, b, e, expBits, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_4096_mod_exp_avx2_64(r, b, e, expBits, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -7692,10 +7775,12 @@ int sp_DhExp_4096(const mp_int* base, const byte* exp, word32 expLen, if (base->used == 1 && base->dp[0] == 2 && m[63] == (sp_digit)-1) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_4096_mod_exp_2_avx2_64(r, e, (int)expLen * 8, m); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_4096_mod_exp_2_avx2_64(r, e, (int)expLen * 8, m); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -7706,10 +7791,12 @@ int sp_DhExp_4096(const mp_int* base, const byte* exp, word32 expLen, { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_4096_mod_exp_avx2_64(r, b, e, (int)expLen * 8, m, 0); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_4096_mod_exp_avx2_64(r, b, e, (int)expLen * 8, m, 0); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -7740,6 +7827,22 @@ int sp_DhExp_4096(const mp_int* base, const byte* exp, word32 expLen, #endif /* WOLFSSL_HAVE_SP_RSA | WOLFSSL_HAVE_SP_DH */ #ifdef WOLFSSL_HAVE_SP_ECC +/* Save only when the cache-resistant table lookup will run, because that is + * the only user of xmm outside the avx2 lane. */ +#ifdef WC_NO_CACHE_RESISTANT +#define SP_ECC_CT_SAVE(ct) 0 +#define SP_ECC_CT_RESTORE(ct) WC_DO_NOTHING +#else +#define SP_ECC_CT_SAVE(ct) ((ct) ? SAVE_VECTOR_REGISTERS2() : 0) +#define SP_ECC_CT_RESTORE(ct) \ + do { \ + if (ct) { \ + RESTORE_VECTOR_REGISTERS();\ + } \ + } \ + while (0) +#endif + #ifndef WOLFSSL_SP_NO_256 /* Point structure to use. */ @@ -10801,7 +10904,13 @@ static int sp_256_ecc_mulmod_4(sp_point_256* r, const sp_point_256* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_256_ecc_mulmod_win_add_sub_4(r, g, k, map, ct, heap); + /* Only the cache-resistant table lookup uses xmm on this lane. */ + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_256_ecc_mulmod_win_add_sub_4(r, g, k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 4 * 5); sp_cache_256_t* cache; @@ -10846,16 +10955,20 @@ static int sp_256_ecc_mulmod_4(sp_point_256* r, const sp_point_256* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_256(g, &cache); - if (cache->cnt == 2) - sp_256_gen_stripe_table_4(g, cache->table, tmp, heap); + err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + sp_ecc_get_cache_256(g, &cache); + if (cache->cnt == 2) + sp_256_gen_stripe_table_4(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_256_ecc_mulmod_win_add_sub_4(r, g, k, map, ct, heap); - } - else { - err = sp_256_ecc_mulmod_stripe_4(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_256_ecc_mulmod_win_add_sub_4(r, g, k, map, ct, heap); + } + else { + err = sp_256_ecc_mulmod_stripe_4(r, g, cache->table, k, + map, ct, heap); + } + SP_ECC_CT_RESTORE(ct); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_256_lock); @@ -11168,7 +11281,13 @@ static int sp_256_ecc_mulmod_avx2_4(sp_point_256* r, const sp_point_256* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_256_ecc_mulmod_win_add_sub_avx2_4(r, g, k, map, ct, heap); + /* The avx2 lane uses vector registers throughout. */ + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_256_ecc_mulmod_win_add_sub_avx2_4(r, g, k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 4 * 5); sp_cache_256_t* cache; @@ -11213,16 +11332,20 @@ static int sp_256_ecc_mulmod_avx2_4(sp_point_256* r, const sp_point_256* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_256(g, &cache); - if (cache->cnt == 2) - sp_256_gen_stripe_table_avx2_4(g, cache->table, tmp, heap); + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_ecc_get_cache_256(g, &cache); + if (cache->cnt == 2) + sp_256_gen_stripe_table_avx2_4(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_256_ecc_mulmod_win_add_sub_avx2_4(r, g, k, map, ct, heap); - } - else { - err = sp_256_ecc_mulmod_stripe_avx2_4(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_256_ecc_mulmod_win_add_sub_avx2_4(r, g, k, map, ct, heap); + } + else { + err = sp_256_ecc_mulmod_stripe_avx2_4(r, g, cache->table, k, + map, ct, heap); + } + RESTORE_VECTOR_REGISTERS(); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_256_lock); @@ -11265,9 +11388,8 @@ int sp_ecc_mulmod_256(const mp_int* km, const ecc_point* gm, ecc_point* r, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_avx2_4(point, point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -11308,6 +11430,7 @@ int sp_ecc_mulmod_add_256(const mp_int* km, const ecc_point* gm, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_256, point, 2, heap, DYNAMIC_TYPE_ECC); @@ -11332,17 +11455,25 @@ int sp_ecc_mulmod_add_256(const mp_int* km, const ecc_point* gm, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_256_ecc_mulmod_avx2_4(point, point, k, 0, 0, heap); + } else #endif err = sp_256_ecc_mulmod_4(point, point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_proj_point_add_avx2_4(point, point, addP, tmp); else #endif @@ -11350,7 +11481,7 @@ int sp_ecc_mulmod_add_256(const mp_int* km, const ecc_point* gm, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_map_avx2_4(point, point, tmp); else #endif @@ -11717,8 +11848,13 @@ static const sp_table_entry_256 p256_table[64] = { static int sp_256_ecc_mulmod_base_4(sp_point_256* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_256_ecc_mulmod_stripe_4(r, &p256_base, p256_table, - k, map, ct, heap); + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_256_ecc_mulmod_stripe_4(r, &p256_base, p256_table, + k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; } #ifdef HAVE_INTEL_AVX2 @@ -11742,8 +11878,13 @@ static int sp_256_ecc_mulmod_base_4(sp_point_256* r, const sp_digit* k, static int sp_256_ecc_mulmod_base_avx2_4(sp_point_256* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_256_ecc_mulmod_stripe_avx2_4(r, &p256_base, p256_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_256_ecc_mulmod_stripe_avx2_4(r, &p256_base, p256_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -23893,8 +24034,13 @@ static int sp_256_ecc_mulmod_add_only_4(sp_point_256* r, const sp_point_256* g, static int sp_256_ecc_mulmod_base_4(sp_point_256* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_256_ecc_mulmod_add_only_4(r, NULL, p256_table, - k, map, ct, heap); + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_256_ecc_mulmod_add_only_4(r, NULL, p256_table, + k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; } #ifdef HAVE_INTEL_AVX2 @@ -24004,8 +24150,13 @@ static int sp_256_ecc_mulmod_add_only_avx2_4(sp_point_256* r, const sp_point_256 static int sp_256_ecc_mulmod_base_avx2_4(sp_point_256* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_256_ecc_mulmod_add_only_avx2_4(r, NULL, p256_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_256_ecc_mulmod_add_only_avx2_4(r, NULL, p256_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -24037,9 +24188,8 @@ int sp_ecc_mulmod_base_256(const mp_int* km, ecc_point* r, int map, void* heap) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_base_avx2_4(point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -24079,6 +24229,7 @@ int sp_ecc_mulmod_base_add_256(const mp_int* km, const ecc_point* am, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_256, point, 2, NULL, DYNAMIC_TYPE_ECC); @@ -24102,17 +24253,25 @@ int sp_ecc_mulmod_base_add_256(const mp_int* km, const ecc_point* am, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_256_ecc_mulmod_base_avx2_4(point, k, 0, 0, heap); + } else #endif err = sp_256_ecc_mulmod_base_4(point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_proj_point_add_avx2_4(point, point, addP, tmp); else #endif @@ -24120,7 +24279,7 @@ int sp_ecc_mulmod_base_add_256(const mp_int* km, const ecc_point* am, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_map_avx2_4(point, point, tmp); else #endif @@ -24251,7 +24410,6 @@ int sp_ecc_make_key_256(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); - int saved_vector_registers = 0; #endif (void)heap; @@ -24272,11 +24430,9 @@ int sp_ecc_make_key_256(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_base_avx2_4(point, k, 1, 1, NULL); + } else #endif err = sp_256_ecc_mulmod_base_4(point, k, 1, 1, NULL); @@ -24285,7 +24441,8 @@ int sp_ecc_make_key_256(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) #ifdef WOLFSSL_VALIDATE_ECC_KEYGEN if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_avx2_4(infinity, point, p256_order, 1, 1, NULL); } @@ -24300,11 +24457,6 @@ int sp_ecc_make_key_256(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) } #endif -#ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) - RESTORE_VECTOR_REGISTERS(); -#endif - if (err == MP_OKAY) { err = sp_256_to_mp(k, priv); } @@ -24485,9 +24637,8 @@ int sp_ecc_secret_gen_256(const mp_int* priv, const ecc_point* pub, byte* out, sp_256_point_from_ecc_point_4(point, pub); #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_avx2_4(point, point, k, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -25240,32 +25391,39 @@ static void sp_256_mont_inv_order_avx2_4(sp_digit* r, const sp_digit* a, static int sp_256_calc_s_4(sp_digit* s, const sp_digit* r, sp_digit* k, sp_digit* x, const sp_digit* e, sp_digit* tmp) { - int err; + int err = MP_OKAY; sp_digit carry; sp_int64 c; sp_digit* kInv = k; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif /* Conv k to Montgomery form (mod order) */ #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) - sp_256_mul_avx2_4(k, k, p256_norm_order); + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + saved_vector_registers = 1; + sp_256_mul_avx2_4(k, k, p256_norm_order); + } + } else #endif sp_256_mul_4(k, k, p256_norm_order); - err = sp_256_mod_4(k, k, p256_order); + if (err == MP_OKAY) + err = sp_256_mod_4(k, k, p256_order); if (err == MP_OKAY) { sp_256_norm_4(k); /* kInv = 1/k mod order */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_mont_inv_order_avx2_4(kInv, k, tmp); else #endif @@ -25274,7 +25432,7 @@ static int sp_256_calc_s_4(sp_digit* s, const sp_digit* r, sp_digit* k, /* s = r * x + e */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_mul_avx2_4(x, x, r); else #endif @@ -25293,7 +25451,7 @@ static int sp_256_calc_s_4(sp_digit* s, const sp_digit* r, sp_digit* k, /* s = s * k^-1 mod order */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_256_mont_mul_order_avx2_4(s, s, kInv); else #endif @@ -25373,10 +25531,8 @@ int sp_ecc_sign_256(const byte* hash, word32 hashLen, WC_RNG* rng, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_base_avx2_4(point, k, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -25638,31 +25794,38 @@ extern void sp_256_mod_inv_avx2_4(sp_digit* r, const sp_digit* a, const sp_digit * @param [in, out] p1 First point to add and holds result. * @param [in] p2 Second point to add. * @param [out] tmp Temporary storage for intermediate numbers. + * @return MP_OKAY, or the vector-register save error when the save is refused. */ -static void sp_256_add_points_4(sp_point_256* p1, const sp_point_256* p2, +static int sp_256_add_points_4(sp_point_256* p1, const sp_point_256* p2, sp_digit* tmp) { + int err = MP_OKAY; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); #endif #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_proj_point_add_avx2_4(p1, p1, p2, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_proj_point_add_avx2_4(p1, p1, p2, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif sp_256_proj_point_add_4(p1, p1, p2, tmp); - if (sp_256_iszero_4(p1->z)) { + if ((err == MP_OKAY) && sp_256_iszero_4(p1->z)) { if (sp_256_iszero_4(p1->x) && sp_256_iszero_4(p1->y)) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_proj_point_dbl_avx2_4(p1, p2, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_proj_point_dbl_avx2_4(p1, p2, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -25677,6 +25840,8 @@ static void sp_256_add_points_4(sp_point_256* p1, const sp_point_256* p2, XMEMCPY(p1->z, p256_norm_mod, sizeof(p256_norm_mod)); } } + + return err; } /* Calculate the verification point: [e/s]G + [r/s]Q @@ -25695,7 +25860,7 @@ static void sp_256_add_points_4(sp_point_256* p1, const sp_point_256* p2, static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, sp_digit* s, sp_digit* u1, sp_digit* u2, sp_digit* tmp, void* heap) { - int err; + int err = MP_OKAY; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); #endif @@ -25703,9 +25868,12 @@ static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, #ifndef WOLFSSL_SP_SMALL #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mod_inv_avx2_4(s, s, p256_order); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mod_inv_avx2_4(s, s, p256_order); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -25713,30 +25881,37 @@ static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, sp_256_mod_inv_4(s, s, p256_order); } #endif /* !WOLFSSL_SP_SMALL */ - { + if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mul_avx2_4(s, s, p256_norm_order); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mul_avx2_4(s, s, p256_norm_order); + RESTORE_VECTOR_REGISTERS(); + } } else #endif { sp_256_mul_4(s, s, p256_norm_order); } - err = sp_256_mod_4(s, s, p256_order); + if (err == MP_OKAY) + err = sp_256_mod_4(s, s, p256_order); } if (err == MP_OKAY) { sp_256_norm_4(s); #ifdef WOLFSSL_SP_SMALL #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mont_inv_order_avx2_4(s, s, tmp); - sp_256_mont_mul_order_avx2_4(u1, u1, s); - sp_256_mont_mul_order_avx2_4(u2, u2, s); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mont_inv_order_avx2_4(s, s, tmp); + sp_256_mont_mul_order_avx2_4(u1, u1, s); + sp_256_mont_mul_order_avx2_4(u2, u2, s); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -25748,10 +25923,13 @@ static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, #else #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mont_mul_order_avx2_4(u1, u1, s); - sp_256_mont_mul_order_avx2_4(u2, u2, s); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mont_mul_order_avx2_4(u1, u1, s); + sp_256_mont_mul_order_avx2_4(u2, u2, s); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -25760,11 +25938,12 @@ static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, sp_256_mont_mul_order_4(u2, u2, s); } #endif /* WOLFSSL_SP_SMALL */ + } + if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_base_avx2_4(p1, u1, 0, 0, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -25778,9 +25957,8 @@ static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_avx2_4(p2, p2, u2, 0, 0, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -25791,7 +25969,7 @@ static int sp_256_calc_vfy_point_4(sp_point_256* p1, sp_point_256* p2, } if (err == MP_OKAY) { - sp_256_add_points_4(p1, p2, tmp); + err = sp_256_add_points_4(p1, p2, tmp); } return err; @@ -25870,22 +26048,22 @@ int sp_ecc_verify_256(const byte* hash, word32 hashLen, const mp_int* pX, /* u1 = r.z'.z' mod prime */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mont_sqr_avx2_4(p1->z, p1->z, p256_mod, p256_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mont_sqr_avx2_4(p1->z, p1->z, p256_mod, p256_mp_mod); + sp_256_mont_mul_avx2_4(u1, u2, p1->z, p256_mod, p256_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif + { sp_256_mont_sqr_4(p1->z, p1->z, p256_mod, p256_mp_mod); -#ifdef HAVE_INTEL_AVX2 - if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mont_mul_avx2_4(u1, u2, p1->z, p256_mod, p256_mp_mod); - RESTORE_VECTOR_REGISTERS(); - } - else -#endif sp_256_mont_mul_4(u1, u2, p1->z, p256_mod, p256_mp_mod); + } + } + if (err == MP_OKAY) { *res = (int)(sp_256_cmp_4(p1->x, u1) == 0); if (*res == 0) { /* Reload r and add order. */ @@ -25906,18 +26084,21 @@ int sp_ecc_verify_256(const byte* hash, word32 hashLen, const mp_int* pX, /* u1 = (r + 1*order).z'.z' mod prime */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mont_mul_avx2_4(u1, u2, p1->z, p256_mod, - p256_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mont_mul_avx2_4(u1, u2, p1->z, p256_mod, + p256_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif { sp_256_mont_mul_4(u1, u2, p1->z, p256_mod, p256_mp_mod); } - *res = (sp_256_cmp_4(p1->x, u1) == 0); + if (err == MP_OKAY) + *res = (sp_256_cmp_4(p1->x, u1) == 0); } } } @@ -26252,9 +26433,8 @@ int sp_ecc_check_key_256(const mp_int* pX, const mp_int* pY, /* Point * order = infinity */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_avx2_4(p, pub, p256_order, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -26271,10 +26451,8 @@ int sp_ecc_check_key_256(const mp_int* pX, const mp_int* pY, /* Base * private = point */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_256_ecc_mulmod_base_avx2_4(p, priv, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -26341,9 +26519,12 @@ int sp_ecc_proj_add_point_256(mp_int* pX, mp_int* pY, mp_int* pZ, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_proj_point_add_avx2_4(p, p, q, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_proj_point_add_avx2_4(p, p, q, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -26400,9 +26581,12 @@ int sp_ecc_proj_dbl_point_256(mp_int* pX, mp_int* pY, mp_int* pZ, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_proj_point_dbl_avx2_4(p, p, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_proj_point_dbl_avx2_4(p, p, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -26456,9 +26640,12 @@ int sp_ecc_map_256(mp_int* pX, mp_int* pY, mp_int* pZ) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_map_avx2_4(p, p, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_map_avx2_4(p, p, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -26504,37 +26691,40 @@ static int sp_256_mont_sqrt_4(sp_digit* y) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* t2 = y ^ 0x2 */ - sp_256_mont_sqr_avx2_4(t2, y, p256_mod, p256_mp_mod); - /* t1 = y ^ 0x3 */ - sp_256_mont_mul_avx2_4(t1, t2, y, p256_mod, p256_mp_mod); - /* t2 = y ^ 0xc */ - sp_256_mont_sqr_n_avx2_4(t2, t1, 2, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xf */ - sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); - /* t2 = y ^ 0xf0 */ - sp_256_mont_sqr_n_avx2_4(t2, t1, 4, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xff */ - sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); - /* t2 = y ^ 0xff00 */ - sp_256_mont_sqr_n_avx2_4(t2, t1, 8, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xffff */ - sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); - /* t2 = y ^ 0xffff0000 */ - sp_256_mont_sqr_n_avx2_4(t2, t1, 16, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xffffffff */ - sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xffffffff00000000 */ - sp_256_mont_sqr_n_avx2_4(t1, t1, 32, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xffffffff00000001 */ - sp_256_mont_mul_avx2_4(t1, t1, y, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xffffffff00000001000000000000000000000000 */ - sp_256_mont_sqr_n_avx2_4(t1, t1, 96, p256_mod, p256_mp_mod); - /* t1 = y ^ 0xffffffff00000001000000000000000000000001 */ - sp_256_mont_mul_avx2_4(t1, t1, y, p256_mod, p256_mp_mod); - sp_256_mont_sqr_n_avx2_4(y, t1, 94, p256_mod, p256_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + /* t2 = y ^ 0x2 */ + sp_256_mont_sqr_avx2_4(t2, y, p256_mod, p256_mp_mod); + /* t1 = y ^ 0x3 */ + sp_256_mont_mul_avx2_4(t1, t2, y, p256_mod, p256_mp_mod); + /* t2 = y ^ 0xc */ + sp_256_mont_sqr_n_avx2_4(t2, t1, 2, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xf */ + sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); + /* t2 = y ^ 0xf0 */ + sp_256_mont_sqr_n_avx2_4(t2, t1, 4, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xff */ + sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); + /* t2 = y ^ 0xff00 */ + sp_256_mont_sqr_n_avx2_4(t2, t1, 8, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xffff */ + sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); + /* t2 = y ^ 0xffff0000 */ + sp_256_mont_sqr_n_avx2_4(t2, t1, 16, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xffffffff */ + sp_256_mont_mul_avx2_4(t1, t1, t2, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xffffffff00000000 */ + sp_256_mont_sqr_n_avx2_4(t1, t1, 32, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xffffffff00000001 */ + sp_256_mont_mul_avx2_4(t1, t1, y, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xffffffff00000001000000000000000000000000 */ + sp_256_mont_sqr_n_avx2_4(t1, t1, 96, p256_mod, p256_mp_mod); + /* t1 = y ^ 0xffffffff00000001000000000000000000000001 */ + sp_256_mont_mul_avx2_4(t1, t1, y, p256_mod, p256_mp_mod); + sp_256_mont_sqr_n_avx2_4(y, t1, 94, p256_mod, p256_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -26606,10 +26796,13 @@ int sp_ecc_uncompress_256(mp_int* xm, int odd, mp_int* ym) /* y = x^3 */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_256_mont_sqr_avx2_4(y, x, p256_mod, p256_mp_mod); - sp_256_mont_mul_avx2_4(y, y, x, p256_mod, p256_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_256_mont_sqr_avx2_4(y, x, p256_mod, p256_mp_mod); + sp_256_mont_mul_avx2_4(y, y, x, p256_mod, p256_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -26617,6 +26810,8 @@ int sp_ecc_uncompress_256(mp_int* xm, int odd, mp_int* ym) sp_256_mont_sqr_4(y, x, p256_mod, p256_mp_mod); sp_256_mont_mul_4(y, y, x, p256_mod, p256_mp_mod); } + } + if (err == MP_OKAY) { /* y = x^3 - 3x */ sp_256_mont_sub_4(y, y, x, p256_mod); sp_256_mont_sub_4(y, y, x, p256_mod); @@ -29832,7 +30027,13 @@ static int sp_384_ecc_mulmod_6(sp_point_384* r, const sp_point_384* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_384_ecc_mulmod_win_add_sub_6(r, g, k, map, ct, heap); + /* Only the cache-resistant table lookup uses xmm on this lane. */ + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_384_ecc_mulmod_win_add_sub_6(r, g, k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 6 * 7); sp_cache_384_t* cache; @@ -29877,16 +30078,20 @@ static int sp_384_ecc_mulmod_6(sp_point_384* r, const sp_point_384* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_384(g, &cache); - if (cache->cnt == 2) - sp_384_gen_stripe_table_6(g, cache->table, tmp, heap); + err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + sp_ecc_get_cache_384(g, &cache); + if (cache->cnt == 2) + sp_384_gen_stripe_table_6(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_384_ecc_mulmod_win_add_sub_6(r, g, k, map, ct, heap); - } - else { - err = sp_384_ecc_mulmod_stripe_6(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_384_ecc_mulmod_win_add_sub_6(r, g, k, map, ct, heap); + } + else { + err = sp_384_ecc_mulmod_stripe_6(r, g, cache->table, k, + map, ct, heap); + } + SP_ECC_CT_RESTORE(ct); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_384_lock); @@ -30202,7 +30407,13 @@ static int sp_384_ecc_mulmod_avx2_6(sp_point_384* r, const sp_point_384* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_384_ecc_mulmod_win_add_sub_avx2_6(r, g, k, map, ct, heap); + /* The avx2 lane uses vector registers throughout. */ + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_384_ecc_mulmod_win_add_sub_avx2_6(r, g, k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 6 * 7); sp_cache_384_t* cache; @@ -30247,16 +30458,20 @@ static int sp_384_ecc_mulmod_avx2_6(sp_point_384* r, const sp_point_384* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_384(g, &cache); - if (cache->cnt == 2) - sp_384_gen_stripe_table_avx2_6(g, cache->table, tmp, heap); + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_ecc_get_cache_384(g, &cache); + if (cache->cnt == 2) + sp_384_gen_stripe_table_avx2_6(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_384_ecc_mulmod_win_add_sub_avx2_6(r, g, k, map, ct, heap); - } - else { - err = sp_384_ecc_mulmod_stripe_avx2_6(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_384_ecc_mulmod_win_add_sub_avx2_6(r, g, k, map, ct, heap); + } + else { + err = sp_384_ecc_mulmod_stripe_avx2_6(r, g, cache->table, k, + map, ct, heap); + } + RESTORE_VECTOR_REGISTERS(); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_384_lock); @@ -30299,9 +30514,8 @@ int sp_ecc_mulmod_384(const mp_int* km, const ecc_point* gm, ecc_point* r, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_avx2_6(point, point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -30342,6 +30556,7 @@ int sp_ecc_mulmod_add_384(const mp_int* km, const ecc_point* gm, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_384, point, 2, heap, DYNAMIC_TYPE_ECC); @@ -30366,17 +30581,25 @@ int sp_ecc_mulmod_add_384(const mp_int* km, const ecc_point* gm, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_384_ecc_mulmod_avx2_6(point, point, k, 0, 0, heap); + } else #endif err = sp_384_ecc_mulmod_6(point, point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_proj_point_add_avx2_6(point, point, addP, tmp); else #endif @@ -30384,7 +30607,7 @@ int sp_ecc_mulmod_add_384(const mp_int* km, const ecc_point* gm, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_map_avx2_6(point, point, tmp); else #endif @@ -30751,8 +30974,13 @@ static const sp_table_entry_384 p384_table[64] = { static int sp_384_ecc_mulmod_base_6(sp_point_384* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_384_ecc_mulmod_stripe_6(r, &p384_base, p384_table, - k, map, ct, heap); + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_384_ecc_mulmod_stripe_6(r, &p384_base, p384_table, + k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; } #ifdef HAVE_INTEL_AVX2 @@ -30776,8 +31004,13 @@ static int sp_384_ecc_mulmod_base_6(sp_point_384* r, const sp_digit* k, static int sp_384_ecc_mulmod_base_avx2_6(sp_point_384* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_384_ecc_mulmod_stripe_avx2_6(r, &p384_base, p384_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_384_ecc_mulmod_stripe_avx2_6(r, &p384_base, p384_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -48741,8 +48974,13 @@ static int sp_384_ecc_mulmod_add_only_6(sp_point_384* r, const sp_point_384* g, static int sp_384_ecc_mulmod_base_6(sp_point_384* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_384_ecc_mulmod_add_only_6(r, NULL, p384_table, - k, map, ct, heap); + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_384_ecc_mulmod_add_only_6(r, NULL, p384_table, + k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; } #ifdef HAVE_INTEL_AVX2 @@ -48852,8 +49090,13 @@ static int sp_384_ecc_mulmod_add_only_avx2_6(sp_point_384* r, const sp_point_384 static int sp_384_ecc_mulmod_base_avx2_6(sp_point_384* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_384_ecc_mulmod_add_only_avx2_6(r, NULL, p384_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_384_ecc_mulmod_add_only_avx2_6(r, NULL, p384_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -48885,9 +49128,8 @@ int sp_ecc_mulmod_base_384(const mp_int* km, ecc_point* r, int map, void* heap) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_base_avx2_6(point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -48927,6 +49169,7 @@ int sp_ecc_mulmod_base_add_384(const mp_int* km, const ecc_point* am, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_384, point, 2, NULL, DYNAMIC_TYPE_ECC); @@ -48950,17 +49193,25 @@ int sp_ecc_mulmod_base_add_384(const mp_int* km, const ecc_point* am, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_384_ecc_mulmod_base_avx2_6(point, k, 0, 0, heap); + } else #endif err = sp_384_ecc_mulmod_base_6(point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_proj_point_add_avx2_6(point, point, addP, tmp); else #endif @@ -48968,7 +49219,7 @@ int sp_ecc_mulmod_base_add_384(const mp_int* km, const ecc_point* am, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_map_avx2_6(point, point, tmp); else #endif @@ -49099,7 +49350,6 @@ int sp_ecc_make_key_384(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); - int saved_vector_registers = 0; #endif (void)heap; @@ -49120,11 +49370,9 @@ int sp_ecc_make_key_384(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_base_avx2_6(point, k, 1, 1, NULL); + } else #endif err = sp_384_ecc_mulmod_base_6(point, k, 1, 1, NULL); @@ -49133,7 +49381,8 @@ int sp_ecc_make_key_384(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) #ifdef WOLFSSL_VALIDATE_ECC_KEYGEN if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_avx2_6(infinity, point, p384_order, 1, 1, NULL); } @@ -49148,11 +49397,6 @@ int sp_ecc_make_key_384(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) } #endif -#ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) - RESTORE_VECTOR_REGISTERS(); -#endif - if (err == MP_OKAY) { err = sp_384_to_mp(k, priv); } @@ -49333,9 +49577,8 @@ int sp_ecc_secret_gen_384(const mp_int* priv, const ecc_point* pub, byte* out, sp_384_point_from_ecc_point_6(point, pub); #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_avx2_6(point, point, k, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -49973,32 +50216,39 @@ static void sp_384_mont_inv_order_avx2_6(sp_digit* r, const sp_digit* a, static int sp_384_calc_s_6(sp_digit* s, const sp_digit* r, sp_digit* k, sp_digit* x, const sp_digit* e, sp_digit* tmp) { - int err; + int err = MP_OKAY; sp_digit carry; sp_int64 c; sp_digit* kInv = k; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif /* Conv k to Montgomery form (mod order) */ #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) - sp_384_mul_avx2_6(k, k, p384_norm_order); + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + saved_vector_registers = 1; + sp_384_mul_avx2_6(k, k, p384_norm_order); + } + } else #endif sp_384_mul_6(k, k, p384_norm_order); - err = sp_384_mod_6(k, k, p384_order); + if (err == MP_OKAY) + err = sp_384_mod_6(k, k, p384_order); if (err == MP_OKAY) { sp_384_norm_6(k); /* kInv = 1/k mod order */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_mont_inv_order_avx2_6(kInv, k, tmp); else #endif @@ -50007,7 +50257,7 @@ static int sp_384_calc_s_6(sp_digit* s, const sp_digit* r, sp_digit* k, /* s = r * x + e */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_mul_avx2_6(x, x, r); else #endif @@ -50026,7 +50276,7 @@ static int sp_384_calc_s_6(sp_digit* s, const sp_digit* r, sp_digit* k, /* s = s * k^-1 mod order */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_384_mont_mul_order_avx2_6(s, s, kInv); else #endif @@ -50106,10 +50356,8 @@ int sp_ecc_sign_384(const byte* hash, word32 hashLen, WC_RNG* rng, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_base_avx2_6(point, k, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -50461,31 +50709,38 @@ static int sp_384_mod_inv_6(sp_digit* r, const sp_digit* a, const sp_digit* m) * @param [in, out] p1 First point to add and holds result. * @param [in] p2 Second point to add. * @param [out] tmp Temporary storage for intermediate numbers. + * @return MP_OKAY, or the vector-register save error when the save is refused. */ -static void sp_384_add_points_6(sp_point_384* p1, const sp_point_384* p2, +static int sp_384_add_points_6(sp_point_384* p1, const sp_point_384* p2, sp_digit* tmp) { + int err = MP_OKAY; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); #endif #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_proj_point_add_avx2_6(p1, p1, p2, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_proj_point_add_avx2_6(p1, p1, p2, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif sp_384_proj_point_add_6(p1, p1, p2, tmp); - if (sp_384_iszero_6(p1->z)) { + if ((err == MP_OKAY) && sp_384_iszero_6(p1->z)) { if (sp_384_iszero_6(p1->x) && sp_384_iszero_6(p1->y)) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_proj_point_dbl_avx2_6(p1, p2, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_proj_point_dbl_avx2_6(p1, p2, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -50502,6 +50757,8 @@ static void sp_384_add_points_6(sp_point_384* p1, const sp_point_384* p2, XMEMCPY(p1->z, p384_norm_mod, sizeof(p384_norm_mod)); } } + + return err; } /* Calculate the verification point: [e/s]G + [r/s]Q @@ -50520,39 +50777,45 @@ static void sp_384_add_points_6(sp_point_384* p1, const sp_point_384* p2, static int sp_384_calc_vfy_point_6(sp_point_384* p1, sp_point_384* p2, sp_digit* s, sp_digit* u1, sp_digit* u2, sp_digit* tmp, void* heap) { - int err; + int err = MP_OKAY; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); #endif #ifndef WOLFSSL_SP_SMALL err = sp_384_mod_inv_6(s, s, p384_order); - if (err == MP_OKAY) #endif /* !WOLFSSL_SP_SMALL */ - { + if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mul_avx2_6(s, s, p384_norm_order); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_mul_avx2_6(s, s, p384_norm_order); + RESTORE_VECTOR_REGISTERS(); + } } else #endif { sp_384_mul_6(s, s, p384_norm_order); } - err = sp_384_mod_6(s, s, p384_order); + if (err == MP_OKAY) + err = sp_384_mod_6(s, s, p384_order); } if (err == MP_OKAY) { sp_384_norm_6(s); #ifdef WOLFSSL_SP_SMALL #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mont_inv_order_avx2_6(s, s, tmp); - sp_384_mont_mul_order_avx2_6(u1, u1, s); - sp_384_mont_mul_order_avx2_6(u2, u2, s); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_mont_inv_order_avx2_6(s, s, tmp); + sp_384_mont_mul_order_avx2_6(u1, u1, s); + sp_384_mont_mul_order_avx2_6(u2, u2, s); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -50564,10 +50827,13 @@ static int sp_384_calc_vfy_point_6(sp_point_384* p1, sp_point_384* p2, #else #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mont_mul_order_avx2_6(u1, u1, s); - sp_384_mont_mul_order_avx2_6(u2, u2, s); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_mont_mul_order_avx2_6(u1, u1, s); + sp_384_mont_mul_order_avx2_6(u2, u2, s); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -50576,11 +50842,12 @@ static int sp_384_calc_vfy_point_6(sp_point_384* p1, sp_point_384* p2, sp_384_mont_mul_order_6(u2, u2, s); } #endif /* WOLFSSL_SP_SMALL */ + } + if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_base_avx2_6(p1, u1, 0, 0, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -50594,9 +50861,8 @@ static int sp_384_calc_vfy_point_6(sp_point_384* p1, sp_point_384* p2, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_avx2_6(p2, p2, u2, 0, 0, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -50607,7 +50873,7 @@ static int sp_384_calc_vfy_point_6(sp_point_384* p1, sp_point_384* p2, } if (err == MP_OKAY) { - sp_384_add_points_6(p1, p2, tmp); + err = sp_384_add_points_6(p1, p2, tmp); } return err; @@ -50686,22 +50952,22 @@ int sp_ecc_verify_384(const byte* hash, word32 hashLen, const mp_int* pX, /* u1 = r.z'.z' mod prime */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mont_sqr_avx2_6(p1->z, p1->z, p384_mod, p384_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_mont_sqr_avx2_6(p1->z, p1->z, p384_mod, p384_mp_mod); + sp_384_mont_mul_avx2_6(u1, u2, p1->z, p384_mod, p384_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif + { sp_384_mont_sqr_6(p1->z, p1->z, p384_mod, p384_mp_mod); -#ifdef HAVE_INTEL_AVX2 - if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mont_mul_avx2_6(u1, u2, p1->z, p384_mod, p384_mp_mod); - RESTORE_VECTOR_REGISTERS(); - } - else -#endif sp_384_mont_mul_6(u1, u2, p1->z, p384_mod, p384_mp_mod); + } + } + if (err == MP_OKAY) { *res = (int)(sp_384_cmp_6(p1->x, u1) == 0); if (*res == 0) { /* Reload r and add order. */ @@ -50722,18 +50988,21 @@ int sp_ecc_verify_384(const byte* hash, word32 hashLen, const mp_int* pX, /* u1 = (r + 1*order).z'.z' mod prime */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mont_mul_avx2_6(u1, u2, p1->z, p384_mod, - p384_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_mont_mul_avx2_6(u1, u2, p1->z, p384_mod, + p384_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif { sp_384_mont_mul_6(u1, u2, p1->z, p384_mod, p384_mp_mod); } - *res = (sp_384_cmp_6(p1->x, u1) == 0); + if (err == MP_OKAY) + *res = (sp_384_cmp_6(p1->x, u1) == 0); } } } @@ -51068,9 +51337,8 @@ int sp_ecc_check_key_384(const mp_int* pX, const mp_int* pY, /* Point * order = infinity */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_avx2_6(p, pub, p384_order, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -51087,10 +51355,8 @@ int sp_ecc_check_key_384(const mp_int* pX, const mp_int* pY, /* Base * private = point */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_384_ecc_mulmod_base_avx2_6(p, priv, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -51157,9 +51423,12 @@ int sp_ecc_proj_add_point_384(mp_int* pX, mp_int* pY, mp_int* pZ, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_proj_point_add_avx2_6(p, p, q, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_proj_point_add_avx2_6(p, p, q, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -51216,9 +51485,12 @@ int sp_ecc_proj_dbl_point_384(mp_int* pX, mp_int* pY, mp_int* pZ, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_proj_point_dbl_avx2_6(p, p, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_proj_point_dbl_avx2_6(p, p, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -51272,9 +51544,12 @@ int sp_ecc_map_384(mp_int* pX, mp_int* pY, mp_int* pZ) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_map_avx2_6(p, p, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_map_avx2_6(p, p, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -51326,62 +51601,65 @@ static int sp_384_mont_sqrt_6(sp_digit* y) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* t2 = y ^ 0x2 */ - sp_384_mont_sqr_avx2_6(t2, y, p384_mod, p384_mp_mod); - /* t1 = y ^ 0x3 */ - sp_384_mont_mul_avx2_6(t1, t2, y, p384_mod, p384_mp_mod); - /* t5 = y ^ 0xc */ - sp_384_mont_sqr_n_avx2_6(t5, t1, 2, p384_mod, p384_mp_mod); - /* t1 = y ^ 0xf */ - sp_384_mont_mul_avx2_6(t1, t1, t5, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x1e */ - sp_384_mont_sqr_avx2_6(t2, t1, p384_mod, p384_mp_mod); - /* t3 = y ^ 0x1f */ - sp_384_mont_mul_avx2_6(t3, t2, y, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x3e0 */ - sp_384_mont_sqr_n_avx2_6(t2, t3, 5, p384_mod, p384_mp_mod); - /* t1 = y ^ 0x3ff */ - sp_384_mont_mul_avx2_6(t1, t3, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x7fe0 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 5, p384_mod, p384_mp_mod); - /* t3 = y ^ 0x7fff */ - sp_384_mont_mul_avx2_6(t3, t3, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x3fff800 */ - sp_384_mont_sqr_n_avx2_6(t2, t3, 15, p384_mod, p384_mp_mod); - /* t4 = y ^ 0x3ffffff */ - sp_384_mont_mul_avx2_6(t4, t3, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0xffffffc000000 */ - sp_384_mont_sqr_n_avx2_6(t2, t4, 30, p384_mod, p384_mp_mod); - /* t1 = y ^ 0xfffffffffffff */ - sp_384_mont_mul_avx2_6(t1, t4, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0xfffffffffffffff000000000000000 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 60, p384_mod, p384_mp_mod); - /* t1 = y ^ 0xffffffffffffffffffffffffffffff */ - sp_384_mont_mul_avx2_6(t1, t1, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0xffffffffffffffffffffffffffffff000000000000000000000000000000 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 120, p384_mod, p384_mp_mod); - /* t1 = y ^ 0xffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff */ - sp_384_mont_mul_avx2_6(t1, t1, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x7fffffffffffffffffffffffffffffffffffffffffffffffffffffffffff8000 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 15, p384_mod, p384_mp_mod); - /* t1 = y ^ 0x7fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff */ - sp_384_mont_mul_avx2_6(t1, t3, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff80000000 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 31, p384_mod, p384_mp_mod); - /* t1 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffff */ - sp_384_mont_mul_avx2_6(t1, t4, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffff0 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 4, p384_mod, p384_mp_mod); - /* t1 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffffc */ - sp_384_mont_mul_avx2_6(t1, t5, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0xfffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffeffffffff0000000000000000 */ - sp_384_mont_sqr_n_avx2_6(t2, t1, 62, p384_mod, p384_mp_mod); - /* t1 = y ^ 0xfffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffeffffffff0000000000000001 */ - sp_384_mont_mul_avx2_6(t1, y, t2, p384_mod, p384_mp_mod); - /* t2 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffffc00000000000000040000000 */ - sp_384_mont_sqr_n_avx2_6(y, t1, 30, p384_mod, p384_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + /* t2 = y ^ 0x2 */ + sp_384_mont_sqr_avx2_6(t2, y, p384_mod, p384_mp_mod); + /* t1 = y ^ 0x3 */ + sp_384_mont_mul_avx2_6(t1, t2, y, p384_mod, p384_mp_mod); + /* t5 = y ^ 0xc */ + sp_384_mont_sqr_n_avx2_6(t5, t1, 2, p384_mod, p384_mp_mod); + /* t1 = y ^ 0xf */ + sp_384_mont_mul_avx2_6(t1, t1, t5, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x1e */ + sp_384_mont_sqr_avx2_6(t2, t1, p384_mod, p384_mp_mod); + /* t3 = y ^ 0x1f */ + sp_384_mont_mul_avx2_6(t3, t2, y, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x3e0 */ + sp_384_mont_sqr_n_avx2_6(t2, t3, 5, p384_mod, p384_mp_mod); + /* t1 = y ^ 0x3ff */ + sp_384_mont_mul_avx2_6(t1, t3, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x7fe0 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 5, p384_mod, p384_mp_mod); + /* t3 = y ^ 0x7fff */ + sp_384_mont_mul_avx2_6(t3, t3, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x3fff800 */ + sp_384_mont_sqr_n_avx2_6(t2, t3, 15, p384_mod, p384_mp_mod); + /* t4 = y ^ 0x3ffffff */ + sp_384_mont_mul_avx2_6(t4, t3, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0xffffffc000000 */ + sp_384_mont_sqr_n_avx2_6(t2, t4, 30, p384_mod, p384_mp_mod); + /* t1 = y ^ 0xfffffffffffff */ + sp_384_mont_mul_avx2_6(t1, t4, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0xfffffffffffffff000000000000000 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 60, p384_mod, p384_mp_mod); + /* t1 = y ^ 0xffffffffffffffffffffffffffffff */ + sp_384_mont_mul_avx2_6(t1, t1, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0xffffffffffffffffffffffffffffff000000000000000000000000000000 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 120, p384_mod, p384_mp_mod); + /* t1 = y ^ 0xffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff */ + sp_384_mont_mul_avx2_6(t1, t1, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x7fffffffffffffffffffffffffffffffffffffffffffffffffffffffffff8000 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 15, p384_mod, p384_mp_mod); + /* t1 = y ^ 0x7fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff */ + sp_384_mont_mul_avx2_6(t1, t3, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff80000000 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 31, p384_mod, p384_mp_mod); + /* t1 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffff */ + sp_384_mont_mul_avx2_6(t1, t4, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffff0 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 4, p384_mod, p384_mp_mod); + /* t1 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffffc */ + sp_384_mont_mul_avx2_6(t1, t5, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0xfffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffeffffffff0000000000000000 */ + sp_384_mont_sqr_n_avx2_6(t2, t1, 62, p384_mod, p384_mp_mod); + /* t1 = y ^ 0xfffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffeffffffff0000000000000001 */ + sp_384_mont_mul_avx2_6(t1, y, t2, p384_mod, p384_mp_mod); + /* t2 = y ^ 0x3fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffbfffffffc00000000000000040000000 */ + sp_384_mont_sqr_n_avx2_6(y, t1, 30, p384_mod, p384_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -51478,10 +51756,13 @@ int sp_ecc_uncompress_384(mp_int* xm, int odd, mp_int* ym) /* y = x^3 */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_384_mont_sqr_avx2_6(y, x, p384_mod, p384_mp_mod); - sp_384_mont_mul_avx2_6(y, y, x, p384_mod, p384_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_384_mont_sqr_avx2_6(y, x, p384_mod, p384_mp_mod); + sp_384_mont_mul_avx2_6(y, y, x, p384_mod, p384_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -51489,6 +51770,8 @@ int sp_ecc_uncompress_384(mp_int* xm, int odd, mp_int* ym) sp_384_mont_sqr_6(y, x, p384_mod, p384_mp_mod); sp_384_mont_mul_6(y, y, x, p384_mod, p384_mp_mod); } + } + if (err == MP_OKAY) { /* y = x^3 - 3x */ sp_384_mont_sub_6(y, y, x, p384_mod); sp_384_mont_sub_6(y, y, x, p384_mod); @@ -54593,7 +54876,13 @@ static int sp_521_ecc_mulmod_9(sp_point_521* r, const sp_point_521* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_521_ecc_mulmod_win_add_sub_9(r, g, k, map, ct, heap); + /* Only the cache-resistant table lookup uses xmm on this lane. */ + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_521_ecc_mulmod_win_add_sub_9(r, g, k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 9 * 6); sp_cache_521_t* cache; @@ -54638,16 +54927,20 @@ static int sp_521_ecc_mulmod_9(sp_point_521* r, const sp_point_521* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_521(g, &cache); - if (cache->cnt == 2) - sp_521_gen_stripe_table_9(g, cache->table, tmp, heap); + err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + sp_ecc_get_cache_521(g, &cache); + if (cache->cnt == 2) + sp_521_gen_stripe_table_9(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_521_ecc_mulmod_win_add_sub_9(r, g, k, map, ct, heap); - } - else { - err = sp_521_ecc_mulmod_stripe_9(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_521_ecc_mulmod_win_add_sub_9(r, g, k, map, ct, heap); + } + else { + err = sp_521_ecc_mulmod_stripe_9(r, g, cache->table, k, + map, ct, heap); + } + SP_ECC_CT_RESTORE(ct); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_521_lock); @@ -54963,7 +55256,13 @@ static int sp_521_ecc_mulmod_avx2_9(sp_point_521* r, const sp_point_521* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_521_ecc_mulmod_win_add_sub_avx2_9(r, g, k, map, ct, heap); + /* The avx2 lane uses vector registers throughout. */ + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_521_ecc_mulmod_win_add_sub_avx2_9(r, g, k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 9 * 6); sp_cache_521_t* cache; @@ -55008,16 +55307,20 @@ static int sp_521_ecc_mulmod_avx2_9(sp_point_521* r, const sp_point_521* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_521(g, &cache); - if (cache->cnt == 2) - sp_521_gen_stripe_table_avx2_9(g, cache->table, tmp, heap); + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_ecc_get_cache_521(g, &cache); + if (cache->cnt == 2) + sp_521_gen_stripe_table_avx2_9(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_521_ecc_mulmod_win_add_sub_avx2_9(r, g, k, map, ct, heap); - } - else { - err = sp_521_ecc_mulmod_stripe_avx2_9(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_521_ecc_mulmod_win_add_sub_avx2_9(r, g, k, map, ct, heap); + } + else { + err = sp_521_ecc_mulmod_stripe_avx2_9(r, g, cache->table, k, + map, ct, heap); + } + RESTORE_VECTOR_REGISTERS(); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_521_lock); @@ -55060,9 +55363,8 @@ int sp_ecc_mulmod_521(const mp_int* km, const ecc_point* gm, ecc_point* r, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_avx2_9(point, point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -55103,6 +55405,7 @@ int sp_ecc_mulmod_add_521(const mp_int* km, const ecc_point* gm, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_521, point, 2, heap, DYNAMIC_TYPE_ECC); @@ -55127,17 +55430,25 @@ int sp_ecc_mulmod_add_521(const mp_int* km, const ecc_point* gm, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_521_ecc_mulmod_avx2_9(point, point, k, 0, 0, heap); + } else #endif err = sp_521_ecc_mulmod_9(point, point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_proj_point_add_avx2_9(point, point, addP, tmp); else #endif @@ -55145,7 +55456,7 @@ int sp_ecc_mulmod_add_521(const mp_int* km, const ecc_point* gm, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_map_avx2_9(point, point, tmp); else #endif @@ -55638,8 +55949,13 @@ static const sp_table_entry_521 p521_table[64] = { static int sp_521_ecc_mulmod_base_9(sp_point_521* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_521_ecc_mulmod_stripe_9(r, &p521_base, p521_table, - k, map, ct, heap); + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_521_ecc_mulmod_stripe_9(r, &p521_base, p521_table, + k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; } #ifdef HAVE_INTEL_AVX2 @@ -55663,8 +55979,13 @@ static int sp_521_ecc_mulmod_base_9(sp_point_521* r, const sp_digit* k, static int sp_521_ecc_mulmod_base_avx2_9(sp_point_521* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_521_ecc_mulmod_stripe_avx2_9(r, &p521_base, p521_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_521_ecc_mulmod_stripe_avx2_9(r, &p521_base, p521_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -89688,8 +90009,13 @@ static int sp_521_ecc_mulmod_add_only_9(sp_point_521* r, const sp_point_521* g, static int sp_521_ecc_mulmod_base_9(sp_point_521* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_521_ecc_mulmod_add_only_9(r, NULL, p521_table, - k, map, ct, heap); + int err = SP_ECC_CT_SAVE(ct); + if (err == 0) { + err = sp_521_ecc_mulmod_add_only_9(r, NULL, p521_table, + k, map, ct, heap); + SP_ECC_CT_RESTORE(ct); + } + return err; } #ifdef HAVE_INTEL_AVX2 @@ -89799,8 +90125,13 @@ static int sp_521_ecc_mulmod_add_only_avx2_9(sp_point_521* r, const sp_point_521 static int sp_521_ecc_mulmod_base_avx2_9(sp_point_521* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_521_ecc_mulmod_add_only_avx2_9(r, NULL, p521_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_521_ecc_mulmod_add_only_avx2_9(r, NULL, p521_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -89832,9 +90163,8 @@ int sp_ecc_mulmod_base_521(const mp_int* km, ecc_point* r, int map, void* heap) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_base_avx2_9(point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -89874,6 +90204,7 @@ int sp_ecc_mulmod_base_add_521(const mp_int* km, const ecc_point* am, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_521, point, 2, NULL, DYNAMIC_TYPE_ECC); @@ -89897,17 +90228,25 @@ int sp_ecc_mulmod_base_add_521(const mp_int* km, const ecc_point* am, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_521_ecc_mulmod_base_avx2_9(point, k, 0, 0, heap); + } else #endif err = sp_521_ecc_mulmod_base_9(point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_proj_point_add_avx2_9(point, point, addP, tmp); else #endif @@ -89915,7 +90254,7 @@ int sp_ecc_mulmod_base_add_521(const mp_int* km, const ecc_point* am, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_map_avx2_9(point, point, tmp); else #endif @@ -90047,7 +90386,6 @@ int sp_ecc_make_key_521(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); - int saved_vector_registers = 0; #endif (void)heap; @@ -90068,11 +90406,9 @@ int sp_ecc_make_key_521(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_base_avx2_9(point, k, 1, 1, NULL); + } else #endif err = sp_521_ecc_mulmod_base_9(point, k, 1, 1, NULL); @@ -90081,7 +90417,8 @@ int sp_ecc_make_key_521(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) #ifdef WOLFSSL_VALIDATE_ECC_KEYGEN if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) { + if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_avx2_9(infinity, point, p521_order, 1, 1, NULL); } @@ -90096,11 +90433,6 @@ int sp_ecc_make_key_521(WC_RNG* rng, mp_int* priv, ecc_point* pub, void* heap) } #endif -#ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) - RESTORE_VECTOR_REGISTERS(); -#endif - if (err == MP_OKAY) { err = sp_521_to_mp(k, priv); } @@ -90281,9 +90613,8 @@ int sp_ecc_secret_gen_521(const mp_int* priv, const ecc_point* pub, byte* out, sp_521_point_from_ecc_point_9(point, pub); #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_avx2_9(point, point, k, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -90976,32 +91307,39 @@ static void sp_521_mont_inv_order_avx2_9(sp_digit* r, const sp_digit* a, static int sp_521_calc_s_9(sp_digit* s, const sp_digit* r, sp_digit* k, sp_digit* x, const sp_digit* e, sp_digit* tmp) { - int err; + int err = MP_OKAY; sp_digit carry; sp_int64 c; sp_digit* kInv = k; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif /* Conv k to Montgomery form (mod order) */ #ifdef HAVE_INTEL_AVX2 + /* CPUID picks the lane; a failed save is an error, not a lane switch. */ if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) - sp_521_mul_avx2_9(k, k, p521_norm_order); + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + saved_vector_registers = 1; + sp_521_mul_avx2_9(k, k, p521_norm_order); + } + } else #endif sp_521_mul_9(k, k, p521_norm_order); - err = sp_521_mod_9(k, k, p521_order); + if (err == MP_OKAY) + err = sp_521_mod_9(k, k, p521_order); if (err == MP_OKAY) { sp_521_norm_9(k); /* kInv = 1/k mod order */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_mont_inv_order_avx2_9(kInv, k, tmp); else #endif @@ -91010,7 +91348,7 @@ static int sp_521_calc_s_9(sp_digit* s, const sp_digit* r, sp_digit* k, /* s = r * x + e */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_mul_avx2_9(x, x, r); else #endif @@ -91029,7 +91367,7 @@ static int sp_521_calc_s_9(sp_digit* s, const sp_digit* r, sp_digit* k, /* s = s * k^-1 mod order */ #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_521_mont_mul_order_avx2_9(s, s, kInv); else #endif @@ -91109,10 +91447,8 @@ int sp_ecc_sign_521(const byte* hash, word32 hashLen, WC_RNG* rng, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_base_avx2_9(point, k, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -91472,31 +91808,38 @@ static int sp_521_mod_inv_9(sp_digit* r, const sp_digit* a, const sp_digit* m) * @param [in, out] p1 First point to add and holds result. * @param [in] p2 Second point to add. * @param [out] tmp Temporary storage for intermediate numbers. + * @return MP_OKAY, or the vector-register save error when the save is refused. */ -static void sp_521_add_points_9(sp_point_521* p1, const sp_point_521* p2, +static int sp_521_add_points_9(sp_point_521* p1, const sp_point_521* p2, sp_digit* tmp) { + int err = MP_OKAY; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); #endif #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_proj_point_add_avx2_9(p1, p1, p2, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_proj_point_add_avx2_9(p1, p1, p2, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif sp_521_proj_point_add_9(p1, p1, p2, tmp); - if (sp_521_iszero_9(p1->z)) { + if ((err == MP_OKAY) && sp_521_iszero_9(p1->z)) { if (sp_521_iszero_9(p1->x) && sp_521_iszero_9(p1->y)) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_proj_point_dbl_avx2_9(p1, p2, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_proj_point_dbl_avx2_9(p1, p2, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -91516,6 +91859,8 @@ static void sp_521_add_points_9(sp_point_521* p1, const sp_point_521* p2, XMEMCPY(p1->z, p521_norm_mod, sizeof(p521_norm_mod)); } } + + return err; } /* Calculate the verification point: [e/s]G + [r/s]Q @@ -91534,39 +91879,45 @@ static void sp_521_add_points_9(sp_point_521* p1, const sp_point_521* p2, static int sp_521_calc_vfy_point_9(sp_point_521* p1, sp_point_521* p2, sp_digit* s, sp_digit* u1, sp_digit* u2, sp_digit* tmp, void* heap) { - int err; + int err = MP_OKAY; #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); #endif #ifndef WOLFSSL_SP_SMALL err = sp_521_mod_inv_9(s, s, p521_order); - if (err == MP_OKAY) #endif /* !WOLFSSL_SP_SMALL */ - { + if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mul_avx2_9(s, s, p521_norm_order); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_mul_avx2_9(s, s, p521_norm_order); + RESTORE_VECTOR_REGISTERS(); + } } else #endif { sp_521_mul_9(s, s, p521_norm_order); } - err = sp_521_mod_9(s, s, p521_order); + if (err == MP_OKAY) + err = sp_521_mod_9(s, s, p521_order); } if (err == MP_OKAY) { sp_521_norm_9(s); #ifdef WOLFSSL_SP_SMALL #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mont_inv_order_avx2_9(s, s, tmp); - sp_521_mont_mul_order_avx2_9(u1, u1, s); - sp_521_mont_mul_order_avx2_9(u2, u2, s); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_mont_inv_order_avx2_9(s, s, tmp); + sp_521_mont_mul_order_avx2_9(u1, u1, s); + sp_521_mont_mul_order_avx2_9(u2, u2, s); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -91578,10 +91929,13 @@ static int sp_521_calc_vfy_point_9(sp_point_521* p1, sp_point_521* p2, #else #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mont_mul_order_avx2_9(u1, u1, s); - sp_521_mont_mul_order_avx2_9(u2, u2, s); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_mont_mul_order_avx2_9(u1, u1, s); + sp_521_mont_mul_order_avx2_9(u2, u2, s); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -91590,11 +91944,12 @@ static int sp_521_calc_vfy_point_9(sp_point_521* p1, sp_point_521* p2, sp_521_mont_mul_order_9(u2, u2, s); } #endif /* WOLFSSL_SP_SMALL */ + } + if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_base_avx2_9(p1, u1, 0, 0, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -91608,9 +91963,8 @@ static int sp_521_calc_vfy_point_9(sp_point_521* p1, sp_point_521* p2, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_avx2_9(p2, p2, u2, 0, 0, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -91621,7 +91975,7 @@ static int sp_521_calc_vfy_point_9(sp_point_521* p1, sp_point_521* p2, } if (err == MP_OKAY) { - sp_521_add_points_9(p1, p2, tmp); + err = sp_521_add_points_9(p1, p2, tmp); } return err; @@ -91704,22 +92058,22 @@ int sp_ecc_verify_521(const byte* hash, word32 hashLen, const mp_int* pX, /* u1 = r.z'.z' mod prime */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mont_sqr_avx2_9(p1->z, p1->z, p521_mod, p521_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_mont_sqr_avx2_9(p1->z, p1->z, p521_mod, p521_mp_mod); + sp_521_mont_mul_avx2_9(u1, u2, p1->z, p521_mod, p521_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif + { sp_521_mont_sqr_9(p1->z, p1->z, p521_mod, p521_mp_mod); -#ifdef HAVE_INTEL_AVX2 - if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mont_mul_avx2_9(u1, u2, p1->z, p521_mod, p521_mp_mod); - RESTORE_VECTOR_REGISTERS(); - } - else -#endif sp_521_mont_mul_9(u1, u2, p1->z, p521_mod, p521_mp_mod); + } + } + if (err == MP_OKAY) { *res = (int)(sp_521_cmp_9(p1->x, u1) == 0); if (*res == 0) { /* Reload r and add order. */ @@ -91740,18 +92094,21 @@ int sp_ecc_verify_521(const byte* hash, word32 hashLen, const mp_int* pX, /* u1 = (r + 1*order).z'.z' mod prime */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mont_mul_avx2_9(u1, u2, p1->z, p521_mod, - p521_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_mont_mul_avx2_9(u1, u2, p1->z, p521_mod, + p521_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif { sp_521_mont_mul_9(u1, u2, p1->z, p521_mod, p521_mp_mod); } - *res = (sp_521_cmp_9(p1->x, u1) == 0); + if (err == MP_OKAY) + *res = (sp_521_cmp_9(p1->x, u1) == 0); } } } @@ -92089,9 +92446,8 @@ int sp_ecc_check_key_521(const mp_int* pX, const mp_int* pY, /* Point * order = infinity */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_avx2_9(p, pub, p521_order, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -92108,10 +92464,8 @@ int sp_ecc_check_key_521(const mp_int* pX, const mp_int* pY, /* Base * private = point */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_521_ecc_mulmod_base_avx2_9(p, priv, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -92178,9 +92532,12 @@ int sp_ecc_proj_add_point_521(mp_int* pX, mp_int* pY, mp_int* pZ, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_proj_point_add_avx2_9(p, p, q, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_proj_point_add_avx2_9(p, p, q, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -92237,9 +92594,12 @@ int sp_ecc_proj_dbl_point_521(mp_int* pX, mp_int* pY, mp_int* pZ, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_proj_point_dbl_avx2_9(p, p, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_proj_point_dbl_avx2_9(p, p, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -92293,9 +92653,12 @@ int sp_ecc_map_521(mp_int* pX, mp_int* pY, mp_int* pZ) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_map_avx2_9(p, p, tmp); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_map_avx2_9(p, p, tmp); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -92346,17 +92709,20 @@ static int sp_521_mont_sqrt_9(sp_digit* y) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int i; - - XMEMCPY(t, y, sizeof(sp_digit) * 9); - for (i=518; i>=0; i--) { - sp_521_mont_sqr_avx2_9(t, t, p521_mod, p521_mp_mod); - if (p521_sqrt_power[i / 64] & ((sp_uint64)1 << (i % 64))) - sp_521_mont_mul_avx2_9(t, t, y, p521_mod, p521_mp_mod); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + int i; + + XMEMCPY(t, y, sizeof(sp_digit) * 9); + for (i=518; i>=0; i--) { + sp_521_mont_sqr_avx2_9(t, t, p521_mod, p521_mp_mod); + if (p521_sqrt_power[i / 64] & ((sp_uint64)1 << (i % 64))) + sp_521_mont_mul_avx2_9(t, t, y, p521_mod, p521_mp_mod); + } + XMEMCPY(y, t, sizeof(sp_digit) * 9); + RESTORE_VECTOR_REGISTERS(); } - XMEMCPY(y, t, sizeof(sp_digit) * 9); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -92408,10 +92774,13 @@ int sp_ecc_uncompress_521(mp_int* xm, int odd, mp_int* ym) /* y = x^3 */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - sp_521_mont_sqr_avx2_9(y, x, p521_mod, p521_mp_mod); - sp_521_mont_mul_avx2_9(y, y, x, p521_mod, p521_mp_mod); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_521_mont_sqr_avx2_9(y, x, p521_mod, p521_mp_mod); + sp_521_mont_mul_avx2_9(y, y, x, p521_mod, p521_mp_mod); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -92419,6 +92788,8 @@ int sp_ecc_uncompress_521(mp_int* xm, int odd, mp_int* ym) sp_521_mont_sqr_9(y, x, p521_mod, p521_mp_mod); sp_521_mont_mul_9(y, y, x, p521_mod, p521_mp_mod); } + } + if (err == MP_OKAY) { /* y = x^3 - 3x */ sp_521_mont_sub_9(y, y, x, p521_mod); sp_521_mont_sub_9(y, y, x, p521_mod); @@ -95678,6 +96049,7 @@ static int sp_1024_ecc_mulmod_16(sp_point_1024* r, const sp_point_1024* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC + /* No xmm table lookups on this lane, so no save. */ return sp_1024_ecc_mulmod_win_add_sub_16(r, g, k, map, ct, heap); #else SP_DECL_VAR(sp_digit, tmp, 2 * 16 * 38); @@ -96031,7 +96403,13 @@ static int sp_1024_ecc_mulmod_avx2_16(sp_point_1024* r, const sp_point_1024* g, const sp_digit* k, int map, int ct, void* heap) { #ifndef FP_ECC - return sp_1024_ecc_mulmod_win_add_sub_avx2_16(r, g, k, map, ct, heap); + /* The avx2 helpers assert a held save, so this lane saves too. */ + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_1024_ecc_mulmod_win_add_sub_avx2_16(r, g, k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; #else SP_DECL_VAR(sp_digit, tmp, 2 * 16 * 38); sp_cache_1024_t* cache; @@ -96076,16 +96454,20 @@ static int sp_1024_ecc_mulmod_avx2_16(sp_point_1024* r, const sp_point_1024* g, #endif /* !SINGLE_THREADED && !HAVE_THREAD_LS */ if (err == MP_OKAY) { - sp_ecc_get_cache_1024(g, &cache); - if (cache->cnt == 2) - sp_1024_gen_stripe_table_avx2_16(g, cache->table, tmp, heap); + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + sp_ecc_get_cache_1024(g, &cache); + if (cache->cnt == 2) + sp_1024_gen_stripe_table_avx2_16(g, cache->table, tmp, heap); - if (cache->cnt < 2) { - err = sp_1024_ecc_mulmod_win_add_sub_avx2_16(r, g, k, map, ct, heap); - } - else { - err = sp_1024_ecc_mulmod_stripe_avx2_16(r, g, cache->table, k, - map, ct, heap); + if (cache->cnt < 2) { + err = sp_1024_ecc_mulmod_win_add_sub_avx2_16(r, g, k, map, ct, heap); + } + else { + err = sp_1024_ecc_mulmod_stripe_avx2_16(r, g, cache->table, k, + map, ct, heap); + } + RESTORE_VECTOR_REGISTERS(); } #if !defined(SINGLE_THREADED) && !defined(HAVE_THREAD_LS) wc_UnLockMutex(&sp_cache_1024_lock); @@ -96128,9 +96510,8 @@ int sp_ecc_mulmod_1024(const mp_int* km, const ecc_point* gm, ecc_point* r, #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_1024_ecc_mulmod_avx2_16(point, point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -99493,6 +99874,7 @@ static const sp_table_entry_1024 p1024_table[256] = { static int sp_1024_ecc_mulmod_base_16(sp_point_1024* r, const sp_digit* k, int map, int ct, void* heap) { + /* No xmm table lookups on this lane, so no save. */ return sp_1024_ecc_mulmod_stripe_16(r, &p1024_base, p1024_table, k, map, ct, heap); } @@ -99518,8 +99900,13 @@ static int sp_1024_ecc_mulmod_base_16(sp_point_1024* r, const sp_digit* k, static int sp_1024_ecc_mulmod_base_avx2_16(sp_point_1024* r, const sp_digit* k, int map, int ct, void* heap) { - return sp_1024_ecc_mulmod_stripe_avx2_16(r, &p1024_base, p1024_table, - k, map, ct, heap); + int err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_1024_ecc_mulmod_stripe_avx2_16(r, &p1024_base, p1024_table, + k, map, ct, heap); + RESTORE_VECTOR_REGISTERS(); + } + return err; } #endif /* HAVE_INTEL_AVX2 */ @@ -99550,9 +99937,8 @@ int sp_ecc_mulmod_base_1024(const mp_int* km, ecc_point* r, int map, void* heap) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_1024_ecc_mulmod_base_avx2_16(point, k, map, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -99592,6 +99978,7 @@ int sp_ecc_mulmod_base_add_1024(const mp_int* km, const ecc_point* am, #ifdef HAVE_INTEL_AVX2 word32 cpuid_flags = cpuid_get_flags(); int saved_vector_registers = 0; + int use_avx2_lane = 0; #endif SP_ALLOC_VAR(sp_point_1024, point, 2, NULL, DYNAMIC_TYPE_ECC); @@ -99615,17 +100002,25 @@ int sp_ecc_mulmod_base_add_1024(const mp_int* km, const ecc_point* am, if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - saved_vector_registers = 1; - if (saved_vector_registers) + IS_INTEL_AVX2(cpuid_flags)) { + use_avx2_lane = 1; err = sp_1024_ecc_mulmod_base_avx2_16(point, k, 0, 0, heap); + } else #endif err = sp_1024_ecc_mulmod_base_16(point, k, 0, 0, heap); } +#ifdef HAVE_INTEL_AVX2 + /* The mulmod saved for itself; this save is for the point operations. */ + if ((err == MP_OKAY) && use_avx2_lane) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) + saved_vector_registers = 1; + } +#endif if (err == MP_OKAY) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_1024_proj_point_add_avx2_16(point, point, addP, tmp); else #endif @@ -99633,7 +100028,7 @@ int sp_ecc_mulmod_base_add_1024(const mp_int* km, const ecc_point* am, if (map) { #ifdef HAVE_INTEL_AVX2 - if (saved_vector_registers) + if (use_avx2_lane) sp_1024_map_avx2_16(point, point, tmp); else #endif @@ -99695,10 +100090,13 @@ int sp_ecc_gen_table_1024(const ecc_point* gm, byte* table, word32* len, sp_1024_point_from_ecc_point_16(point, gm); #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_1024_gen_stripe_table_avx2_16(point, - (sp_table_entry_1024*)table, t, heap); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_1024_gen_stripe_table_avx2_16(point, + (sp_table_entry_1024*)table, t, heap); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -99784,10 +100182,13 @@ int sp_ecc_mulmod_table_1024(const mp_int* km, const ecc_point* gm, byte* table, #ifndef WOLFSSL_SP_SMALL #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_1024_ecc_mulmod_stripe_avx2_16(point, point, - (const sp_table_entry_1024*)table, k, map, 0, heap); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_1024_ecc_mulmod_stripe_avx2_16(point, point, + (const sp_table_entry_1024*)table, k, map, 0, heap); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -101876,9 +102277,12 @@ int sp_ModExp_Fp_star_1024(const mp_int* base, mp_int* exp, mp_int* res) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_ModExp_Fp_star_avx2_1024(base, exp, res); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_ModExp_Fp_star_avx2_1024(base, exp, res); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -103484,9 +103888,12 @@ int sp_Pairing_1024(const ecc_point* pm, const ecc_point* qm, mp_int* res) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_Pairing_avx2_1024(pm, qm, res); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_Pairing_avx2_1024(pm, qm, res); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -104618,9 +105025,12 @@ int sp_Pairing_gen_precomp_1024(const ecc_point* pm, byte* table, word32* len) #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_Pairing_gen_precomp_avx2_1024(pm, table, len); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_Pairing_gen_precomp_avx2_1024(pm, table, len); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -104656,9 +105066,12 @@ int sp_Pairing_precomp_1024(const ecc_point* pm, const ecc_point* qm, mp_int* re #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - err = sp_Pairing_precomp_avx2_1024(pm, qm, res, table, len); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + err = SAVE_VECTOR_REGISTERS2(); + if (err == 0) { + err = sp_Pairing_precomp_avx2_1024(pm, qm, res, table, len); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -104857,9 +105270,8 @@ int sp_ecc_check_key_1024(const mp_int* pX, const mp_int* pY, /* Point * order = infinity */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_1024_ecc_mulmod_avx2_16(p, pub, p1024_order, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif @@ -104876,10 +105288,8 @@ int sp_ecc_check_key_1024(const mp_int* pX, const mp_int* pY, /* Base * private = point */ #ifdef HAVE_INTEL_AVX2 if (IS_INTEL_BMI2(cpuid_flags) && IS_INTEL_ADX(cpuid_flags) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { err = sp_1024_ecc_mulmod_base_avx2_16(p, priv, 1, 1, heap); - RESTORE_VECTOR_REGISTERS(); } else #endif diff --git a/wolfcrypt/src/wc_lms_impl.c b/wolfcrypt/src/wc_lms_impl.c index 74535332084..b8d052d9cba 100644 --- a/wolfcrypt/src/wc_lms_impl.c +++ b/wolfcrypt/src/wc_lms_impl.c @@ -2269,22 +2269,26 @@ static int wc_lmots_compute_y_from_seed(LmsState* state, const byte* seed, if (ret == 0) { int lanes = LMS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* Chain i runs from x[i] for a[i] steps, into y[i]. */ + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + /* Chain i runs from x[i] for a[i] steps, into y[i]. */ #ifdef WC_LMS_N_WAY_FUSED - if (LMS_N_WAY_FUSED(params, lanes)) { - ret = wc_lmots_n_way_chains_fused(state, LMS_N_WAY_TO_A, seed, - NULL, a, 0, y, lanes); - } - else + if (LMS_N_WAY_FUSED(params, lanes)) { + ret = wc_lmots_n_way_chains_fused(state, LMS_N_WAY_TO_A, + seed, NULL, a, 0, y, lanes); + } + else #endif - { - ret = wc_lmots_n_way_chains(state, LMS_N_WAY_TO_A, seed, - NULL, a, - 0, y, lanes); + { + ret = wc_lmots_n_way_chains(state, LMS_N_WAY_TO_A, seed, + NULL, a, + 0, y, lanes); + } + RESTORE_VECTOR_REGISTERS(); + i = params->p; } - RESTORE_VECTOR_REGISTERS(); - i = params->p; } } #endif @@ -2472,22 +2476,27 @@ static int wc_lmots_compute_kc_from_sig(LmsState* state, const byte* msg, if (ret == 0) { int lanes = LMS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* Chain i resumes at a[i] and runs to 2^w - 1; each result - * is hashed into Kc in chain order. */ + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + /* Chain i resumes at a[i] and runs to 2^w - 1; each result + * is hashed into Kc in chain order. */ #ifdef WC_LMS_N_WAY_FUSED - if (LMS_N_WAY_FUSED(params, lanes)) { - ret = wc_lmots_n_way_chains_fused(state, LMS_N_WAY_FROM_A, - NULL, sig_y, a, (word16)max, NULL, lanes); - } - else + if (LMS_N_WAY_FUSED(params, lanes)) { + ret = wc_lmots_n_way_chains_fused(state, + LMS_N_WAY_FROM_A, NULL, sig_y, a, (word16)max, + NULL, lanes); + } + else #endif - { - ret = wc_lmots_n_way_chains(state, LMS_N_WAY_FROM_A, NULL, - sig_y, a, (word16)max, NULL, lanes); + { + ret = wc_lmots_n_way_chains(state, LMS_N_WAY_FROM_A, + NULL, sig_y, a, (word16)max, NULL, lanes); + } + RESTORE_VECTOR_REGISTERS(); + i = params->p; } - RESTORE_VECTOR_REGISTERS(); - i = params->p; } } #endif @@ -2565,22 +2574,27 @@ static int wc_lmots_compute_kc_from_sig(LmsState* state, const byte* msg, if (ret == 0) { int lanes = LMS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* Chain i resumes at a[i] and runs to 2^w - 1; each result - * is hashed into Kc in chain order. */ + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + /* Chain i resumes at a[i] and runs to 2^w - 1; each result + * is hashed into Kc in chain order. */ #ifdef WC_LMS_N_WAY_FUSED - if (LMS_N_WAY_FUSED(params, lanes)) { - ret = wc_lmots_n_way_chains_fused(state, LMS_N_WAY_FROM_A, - NULL, sig_y, a, (word16)max, NULL, lanes); - } - else + if (LMS_N_WAY_FUSED(params, lanes)) { + ret = wc_lmots_n_way_chains_fused(state, + LMS_N_WAY_FROM_A, NULL, sig_y, a, (word16)max, + NULL, lanes); + } + else #endif - { - ret = wc_lmots_n_way_chains(state, LMS_N_WAY_FROM_A, NULL, - sig_y, a, (word16)max, NULL, lanes); + { + ret = wc_lmots_n_way_chains(state, LMS_N_WAY_FROM_A, + NULL, sig_y, a, (word16)max, NULL, lanes); + } + RESTORE_VECTOR_REGISTERS(); + i = params->p; } - RESTORE_VECTOR_REGISTERS(); - i = params->p; } } #endif @@ -2715,25 +2729,29 @@ static int wc_lmots_make_public_hash(LmsState* state, const byte* seed, byte* k) if (ret == 0) { int lanes = LMS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* Every chain runs the full 2^w - 1 iterations, so the - * lanes stay in step and the batch needs no scheduling. */ + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + /* Every chain runs the full 2^w - 1 iterations, so the + * lanes stay in step and the batch needs no scheduling. */ #ifdef WC_LMS_N_WAY_FUSED - /* The kernels lay out a 32-byte tmp: eight state words in - * and the 0x80 in W13. The 24-byte parameter sets put the - * padding elsewhere, so they keep the general path. */ - if (LMS_N_WAY_FUSED(params, lanes)) { - ret = wc_lmots_n_way_pub_chains_fused(state, seed, - (word16)max, lanes); - } - else + /* The kernels lay out a 32-byte tmp: eight state words in + * and the 0x80 in W13. The 24-byte parameter sets put the + * padding elsewhere, so they keep the general path. */ + if (LMS_N_WAY_FUSED(params, lanes)) { + ret = wc_lmots_n_way_pub_chains_fused(state, seed, + (word16)max, lanes); + } + else #endif - { - ret = wc_lmots_n_way_pub_chains(state, seed, (word16)max, - lanes); + { + ret = wc_lmots_n_way_pub_chains(state, seed, + (word16)max, lanes); + } + RESTORE_VECTOR_REGISTERS(); + i = params->p; } - RESTORE_VECTOR_REGISTERS(); - i = params->p; } } #endif @@ -2798,25 +2816,29 @@ static int wc_lmots_make_public_hash(LmsState* state, const byte* seed, byte* k) if (ret == 0) { int lanes = LMS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - /* Every chain runs the full 2^w - 1 iterations, so the - * lanes stay in step and the batch needs no scheduling. */ + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + /* Every chain runs the full 2^w - 1 iterations, so the + * lanes stay in step and the batch needs no scheduling. */ #ifdef WC_LMS_N_WAY_FUSED - /* The kernels lay out a 32-byte tmp: eight state words in - * and the 0x80 in W13. The 24-byte parameter sets put the - * padding elsewhere, so they keep the general path. */ - if (LMS_N_WAY_FUSED(params, lanes)) { - ret = wc_lmots_n_way_pub_chains_fused(state, seed, - (word16)max, lanes); - } - else + /* The kernels lay out a 32-byte tmp: eight state words in + * and the 0x80 in W13. The 24-byte parameter sets put the + * padding elsewhere, so they keep the general path. */ + if (LMS_N_WAY_FUSED(params, lanes)) { + ret = wc_lmots_n_way_pub_chains_fused(state, seed, + (word16)max, lanes); + } + else #endif - { - ret = wc_lmots_n_way_pub_chains(state, seed, (word16)max, - lanes); + { + ret = wc_lmots_n_way_pub_chains(state, seed, + (word16)max, lanes); + } + RESTORE_VECTOR_REGISTERS(); + i = params->p; } - RESTORE_VECTOR_REGISTERS(); - i = params->p; } } #endif diff --git a/wolfcrypt/src/wc_mldsa.c b/wolfcrypt/src/wc_mldsa.c index 8fade7cfb5f..55d63a4a6a5 100644 --- a/wolfcrypt/src/wc_mldsa.c +++ b/wolfcrypt/src/wc_mldsa.c @@ -209,15 +209,16 @@ static cpuid_flags_t cpuid_flags = WC_CPUID_INITIALIZER; /* AVX2 NTT/invNTT flavor selection: the non-full AVX2 NTT/invNTT keep the * NTT-domain coefficients in a permuted (lane-interleaved) order that only * the non-full AVX2 consumers understand, whereas the full variants and the - * C implementations all use the standard order. With WC_C_DYNAMIC_FALLBACK, - * SAVE_VECTOR_REGISTERS2() can fail on any call, so NTT-domain data at rest - * (cached s1/s2/t0 vectors, the challenge polynomial, etc.) can be produced - * and consumed by differently-dispatched calls, and its representation must - * be dispatch-invariant, i.e. standard order. Without WC_C_DYNAMIC_FALLBACK, - * SAVE_VECTOR_REGISTERS2() cannot fail intermittently (fuzzing without - * fallback is an unsupported contradiction, and kernel-mode intelasm builds - * always define WC_C_DYNAMIC_FALLBACK), so dispatch is invariant and the - * slightly faster (~2%/~4% on NTT/invNTT) permuted-order variants are safe. + * C implementations all use the standard order. A refused save in this file + * is an error, never a switch to another lane, so dispatch no longer changes + * from one call to the next on its own. The choice still keys off + * WC_C_DYNAMIC_FALLBACK, which is where other files still switch lanes at run + * time, and the standard-order variants stay readable whatever ran before. + * Without it the slightly faster (~2%/~4% on NTT/invNTT) permuted-order + * variants are used; they assume the lane in force when NTT-domain data at + * rest (cached s1/s2/t0 vectors, the challenge polynomial) was produced is the + * lane that reads it back, so cpuid_set_flag()/cpuid_clear_flag() must not be + * used to change lanes while such a key is live. * Both pipelines yield bit-identical end results. */ #ifdef WC_C_DYNAMIC_FALLBACK #define MLDSA_NTT_AVX512(r) wc_mldsa_ntt_full_1p_avx512(r) @@ -524,8 +525,10 @@ static int mldsa_shake256(wc_Shake* shake256, const byte* data, dataLen -= WC_SHA3_256_COUNT * 8; data += WC_SHA3_256_COUNT * 8; #ifndef WC_SHA3_NO_ASM - if (SHA3_USE_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -540,8 +543,10 @@ static int mldsa_shake256(wc_Shake* shake256, const byte* data, if (dataLen >= WC_SHA3_256_COUNT * 8) { #ifndef WC_SHA3_NO_ASM word32 n = dataLen / (WC_SHA3_256_COUNT * 8); - if (SHA3_USE_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_n_avx2(state, data, n, WC_SHA3_256_COUNT * 8); RESTORE_VECTOR_REGISTERS(); n *= WC_SHA3_256_COUNT * 8; @@ -579,7 +584,10 @@ static int mldsa_shake256(wc_Shake* shake256, const byte* data, } #ifndef WC_SHA3_NO_ASM - if (SHA3_USE_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -649,8 +657,10 @@ static int mldsa_hash256(wc_Shake* shake256, const byte* data1, data2Len -= WC_SHA3_256_COUNT * 8 - data1Len; data2 += WC_SHA3_256_COUNT * 8 - data1Len; #ifndef WC_SHA3_NO_ASM - if (SHA3_USE_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -666,7 +676,10 @@ static int mldsa_hash256(wc_Shake* shake256, const byte* data1, if (data2Len >= WC_SHA3_256_COUNT * 8) { #ifndef WC_SHA3_NO_ASM word32 n = data2Len / (WC_SHA3_256_COUNT * 8); - if (SHA3_USE_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_n_avx2(state, data2, n, WC_SHA3_256_COUNT * 8); RESTORE_VECTOR_REGISTERS(); n *= WC_SHA3_256_COUNT * 8; @@ -708,7 +721,10 @@ static int mldsa_hash256(wc_Shake* shake256, const byte* data1, } #ifndef WC_SHA3_NO_ASM - if (SHA3_USE_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -1006,8 +1022,10 @@ static int mldsa_squeeze256(wc_Shake* shake256, const byte* in, for (; outBlocks > 0; outBlocks--) { #ifndef WC_SHA3_NO_ASM - if (SHA3_USE_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -1157,17 +1175,21 @@ static void mldsa_vec_encode_eta_bits_c(const sword32* s, byte d, byte eta, * @param [in] eta Range specifier of each value. * @param [out] p Buffer to encode into. */ -static void mldsa_vec_encode_eta_bits(const sword32* s, byte d, byte eta, +static int mldsa_vec_encode_eta_bits(const sword32* s, byte d, byte eta, byte* p) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* eta = 4 packs one polynomial per call - vpmovqb needs no lane pairing. * eta = 2 still pairs, so an odd trailing polynomial is left to the C * code: the AVX2 entry points take a whole vector, not one polynomial. */ - if (USE_INTEL_AVX512(cpuid_flags) && (eta == MLDSA_ETA_4) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags) && (eta == MLDSA_ETA_4)) { unsigned int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; for (i = 0; i < d; i++) { wc_mldsa_encode_eta_4_avx512(s, p); s += MLDSA_N; @@ -1175,14 +1197,14 @@ static void mldsa_vec_encode_eta_bits(const sword32* s, byte d, byte eta, } RESTORE_VECTOR_REGISTERS(); } - /* Note the eta check: this arm is also reachable for eta == 4, when the - * first arm's SAVE_VECTOR_REGISTERS2() fails (e.g. under - * DEBUG_VECTOR_REGISTER_ACCESS_FUZZING), and must not consume that flow. - */ else if (USE_INTEL_AVX512(cpuid_flags) && (eta == MLDSA_ETA_2) && - ((d & 1) == 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { + ((d & 1) == 0)) { unsigned int i; unsigned int e = MLDSA_ETA_2_BITS * MLDSA_N / 8; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; for (i = 0; i < d; i += 2) { wc_mldsa_encode_eta_2_x2_avx512(s, s + MLDSA_N, p, p + e); s += 2 * MLDSA_N; @@ -1192,7 +1214,10 @@ static void mldsa_vec_encode_eta_bits(const sword32* s, byte d, byte eta, } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; #if !defined(WOLFSSL_NO_ML_DSA_44) || !defined(WOLFSSL_NO_ML_DSA_87) /* -2..2 */ if (eta == MLDSA_ETA_2) { @@ -1211,6 +1236,7 @@ static void mldsa_vec_encode_eta_bits(const sword32* s, byte d, byte eta, { mldsa_vec_encode_eta_bits_c(s, d, eta, p); } + return ret; } #endif /* !WOLFSSL_MLDSA_NO_MAKE_KEY */ @@ -1264,10 +1290,14 @@ static void mldsa_decode_eta_2_bits_c(const byte* p, sword32* s) * @param [in] p Buffer of data to decode. * @param [in] s Vector of decoded polynomials. */ -static void mldsa_decode_eta_2_bits(const byte* p, sword32* s) +static int mldsa_decode_eta_2_bits(const byte* p, sword32* s) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_decode_eta_2_avx2(p, s); RESTORE_VECTOR_REGISTERS(); } @@ -1276,6 +1306,7 @@ static void mldsa_decode_eta_2_bits(const byte* p, sword32* s) { mldsa_decode_eta_2_bits_c(p, s); } + return ret; } #endif #ifndef WOLFSSL_NO_ML_DSA_65 @@ -1334,10 +1365,14 @@ static void mldsa_decode_eta_4_bits_c(const byte* p, sword32* s) * @param [in] p Buffer of data to decode. * @param [in] s Vector of decoded polynomials. */ -static void mldsa_decode_eta_4_bits(const byte* p, sword32* s) +static int mldsa_decode_eta_4_bits(const byte* p, sword32* s) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_decode_eta_4_avx2(p, s); RESTORE_VECTOR_REGISTERS(); } @@ -1346,6 +1381,7 @@ static void mldsa_decode_eta_4_bits(const byte* p, sword32* s) { mldsa_decode_eta_4_bits_c(p, s); } + return ret; } #endif @@ -1374,9 +1410,10 @@ static void mldsa_decode_eta_4_bits(const byte* p, sword32* s) * @param [in] s Vector of decoded polynomials. * @param [in] d Dimension of vector. */ -static void mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, +static int mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, byte d) { + int ret = 0; unsigned int i; #if !defined(WOLFSSL_NO_ML_DSA_44) || !defined(WOLFSSL_NO_ML_DSA_87) @@ -1386,7 +1423,10 @@ static void mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* Two polynomials per call; an odd trailing one falls through. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i + 1 < d; i += 2) { wc_mldsa_decode_eta_2_x2_avx512(p, p + e, s, s + MLDSA_N); p += 2 * e; @@ -1396,8 +1436,8 @@ static void mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, } #endif /* Step 5 or 8: For each polynomial of vector */ - for (; i < d; i++) { - mldsa_decode_eta_2_bits(p, s); + for (; (ret == 0) && (i < d); i++) { + ret = mldsa_decode_eta_2_bits(p, s); /* Move to next place to decode from. */ p += e; /* Next polynomial. */ @@ -1411,7 +1451,10 @@ static void mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, unsigned int e = MLDSA_N / 2; i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i + 1 < d; i += 2) { wc_mldsa_decode_eta_4_x2_avx512(p, p + e, s, s + MLDSA_N); p += 2 * e; @@ -1421,8 +1464,8 @@ static void mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, } #endif /* Step 5 or 8: For each polynomial of vector */ - for (; i < d; i++) { - mldsa_decode_eta_4_bits(p, s); + for (; (ret == 0) && (i < d); i++) { + ret = mldsa_decode_eta_4_bits(p, s); /* Move to next place to decode from. */ p += e; /* Next polynomial. */ @@ -1430,6 +1473,7 @@ static void mldsa_vec_decode_eta_bits(const byte* p, byte eta, sword32* s, } } #endif + return ret; } #endif #endif /* !WOLFSSL_MLDSA_NO_SIGN || WOLFSSL_MLDSA_CHECK_KEY */ @@ -1567,19 +1611,23 @@ static void mldsa_vec_encode_t0_t1_c(const sword32* t, byte d, byte* t0, * @param [out] t0 Buffer to encode bottom part of value of t into. * @param [out] t1 Buffer to encode top part of value of t into. */ -static void mldsa_vec_encode_t0_t1(const sword32* t, byte d, byte* t0, +static int mldsa_vec_encode_t0_t1(const sword32* t, byte d, byte* t0, byte* t1) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512_VBMI /* vpermb gathers a whole register's encoded bytes, so this needs no * polynomial pairing and takes any dimension. */ - if (IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX512_VBMI(cpuid_flags)) { unsigned int i; unsigned int e0 = MLDSA_D * MLDSA_N / 8; unsigned int e1 = MLDSA_U * MLDSA_N / 8; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; for (i = 0; i < d; i++) { wc_mldsa_encode_t0_t1_avx512_vbmi(t, t0, t1); t += MLDSA_N; @@ -1592,11 +1640,14 @@ static void mldsa_vec_encode_t0_t1(const sword32* t, byte d, byte* t0, #endif /* Two polynomials per call; the AVX2 entry takes a whole vector, so an * odd dimension stays on it. */ - if (USE_INTEL_AVX512(cpuid_flags) && ((d & 1) == 0) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags) && ((d & 1) == 0)) { unsigned int i; unsigned int e0 = MLDSA_D * MLDSA_N / 8; unsigned int e1 = MLDSA_U * MLDSA_N / 8; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; for (i = 0; i < d; i += 2) { wc_mldsa_encode_t0_t1_x2_avx512(t, t + MLDSA_N, t0, t0 + e0, t1, t1 + e1); @@ -1608,7 +1659,10 @@ static void mldsa_vec_encode_t0_t1(const sword32* t, byte d, byte* t0, } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_vec_encode_t0_t1_avx2(t, d, t0, t1); RESTORE_VECTOR_REGISTERS(); } @@ -1617,6 +1671,7 @@ static void mldsa_vec_encode_t0_t1(const sword32* t, byte d, byte* t0, { mldsa_vec_encode_t0_t1_c(t, d, t0, t1); } + return ret; } #endif /* !WOLFSSL_MLDSA_NO_MAKE_KEY */ @@ -1703,10 +1758,14 @@ static void mldsa_decode_t0_c(const byte* t0, sword32* t) * @param [in] t0 Encoded values of t0. * @param [out] t Vector of polynomials. */ -static void mldsa_decode_t0(const byte* t0, sword32* t) +static int mldsa_decode_t0(const byte* t0, sword32* t) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_decode_t0_avx2(t0, t); RESTORE_VECTOR_REGISTERS(); } @@ -1715,6 +1774,7 @@ static void mldsa_decode_t0(const byte* t0, sword32* t) { mldsa_decode_t0_c(t0, t); } + return ret; } #if defined(WOLFSSL_MLDSA_CHECK_KEY) || \ @@ -1734,14 +1794,18 @@ static void mldsa_decode_t0(const byte* t0, sword32* t) * @param [in] d Dimensions of vector t0. * @param [out] t Vector of polynomials. */ -static void mldsa_vec_decode_t0(const byte* t0, byte d, sword32* t) +static int mldsa_vec_decode_t0(const byte* t0, byte d, sword32* t) { + int ret = 0; unsigned int i; i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* Two polynomials per call; an odd trailing one falls through. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i + 1 < d; i += 2) { wc_mldsa_decode_t0_x2_avx512(t0, t0 + MLDSA_D * MLDSA_N / 8, t, t + MLDSA_N); @@ -1752,12 +1816,13 @@ static void mldsa_vec_decode_t0(const byte* t0, byte d, sword32* t) } #endif /* Step 11. For each polynomial of vector. */ - for (; i < d; i++) { - mldsa_decode_t0(t0, t); + for (; (ret == 0) && (i < d); i++) { + ret = mldsa_decode_t0(t0, t); t0 += MLDSA_D * MLDSA_N / 8; /* Next polynomial. */ t += MLDSA_N; } + return ret; } #endif #endif /* !WOLFSSL_MLDSA_NO_SIGN || WOLFSSL_MLDSA_CHECK_KEY */ @@ -1847,10 +1912,14 @@ static void mldsa_decode_t1_c(const byte* t1, sword32* t) * @param [in] t1 Encoded values of t1. * @param [out] t Polynomials. */ -static void mldsa_decode_t1(const byte* t1, sword32* t) +static int mldsa_decode_t1(const byte* t1, sword32* t) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_decode_t1_avx2(t1, t); RESTORE_VECTOR_REGISTERS(); } @@ -1859,6 +1928,7 @@ static void mldsa_decode_t1(const byte* t1, sword32* t) { mldsa_decode_t1_c(t1, t); } + return ret; } #endif @@ -1878,14 +1948,18 @@ static void mldsa_decode_t1(const byte* t1, sword32* t) * @param [in] d Dimensions of vector t1. * @param [out] t Vector of polynomials. */ -static void mldsa_vec_decode_t1(const byte* t1, byte d, sword32* t) +static int mldsa_vec_decode_t1(const byte* t1, byte d, sword32* t) { + int ret = 0; unsigned int i; i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* Two polynomials per call; an odd trailing one falls through. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i + 1 < d; i += 2) { wc_mldsa_decode_t1_x2_avx512(t1, t1 + MLDSA_U * MLDSA_N / 8, t, t + MLDSA_N); @@ -1896,12 +1970,13 @@ static void mldsa_vec_decode_t1(const byte* t1, byte d, sword32* t) } #endif /* Step 3. For each polynomial of vector. */ - for (; i < d; i++) { - mldsa_decode_t1(t1, t); + for (; (ret == 0) && (i < d); i++) { + ret = mldsa_decode_t1(t1, t); /* Next polynomial. */ t1 += MLDSA_U * MLDSA_N / 8; t += MLDSA_N; } + return ret; } #endif @@ -1960,17 +2035,24 @@ static void mldsa_encode_gamma1_17_bits_c(const sword32* z, byte* s) * @param [in] z Polynomial to encode. * @param [out] s Buffer to encode into. */ -static void mldsa_encode_gamma1_17_bits(const sword32* z, byte* s) +static int mldsa_encode_gamma1_17_bits(const sword32* z, byte* s) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_encode_gamma1_17_avx512(z, s); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_encode_gamma1_17_avx2(z, s); RESTORE_VECTOR_REGISTERS(); } @@ -1979,6 +2061,7 @@ static void mldsa_encode_gamma1_17_bits(const sword32* z, byte* s) { mldsa_encode_gamma1_17_bits_c(z, s); } + return ret; } #endif #if !defined(WOLFSSL_NO_ML_DSA_65) || !defined(WOLFSSL_NO_ML_DSA_87) @@ -2036,17 +2119,24 @@ static void mldsa_encode_gamma1_19_bits_c(const sword32* z, byte* s) * @param [in] z Polynomial to encode. * @param [out] s Buffer to encode into. */ -static void mldsa_encode_gamma1_19_bits(const sword32* z, byte* s) +static int mldsa_encode_gamma1_19_bits(const sword32* z, byte* s) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_encode_gamma1_19_avx512(z, s); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_encode_gamma1_19_avx2(z, s); RESTORE_VECTOR_REGISTERS(); } @@ -2055,6 +2145,7 @@ static void mldsa_encode_gamma1_19_bits(const sword32* z, byte* s) { mldsa_encode_gamma1_19_bits_c(z, s); } + return ret; } #endif @@ -2073,9 +2164,10 @@ static void mldsa_encode_gamma1_19_bits(const sword32* z, byte* s) * @param [in] bits Number of bits used in encoding - GAMMA1 bits. * @param [out] s Buffer to encode into. */ -static void mldsa_vec_encode_gamma1(const sword32* z, byte l, int bits, +static int mldsa_vec_encode_gamma1(const sword32* z, byte l, int bits, byte* s) { + int ret = 0; unsigned int i; (void)l; @@ -2086,8 +2178,8 @@ static void mldsa_vec_encode_gamma1(const sword32* z, byte l, int bits, * polynomial per call, so the dispatch is inside the per-polynomial * function. */ /* Step 2. For each polynomial of vector. */ - for (i = 0; i < PARAMS_ML_DSA_44_L; i++) { - mldsa_encode_gamma1_17_bits(z, s); + for (i = 0; (ret == 0) && (i < PARAMS_ML_DSA_44_L); i++) { + ret = mldsa_encode_gamma1_17_bits(z, s); /* Move to next place to encode to. */ s += MLDSA_GAMMA1_17_ENC_BITS / 2 * MLDSA_N / 4; /* Next polynomial. */ @@ -2099,8 +2191,8 @@ static void mldsa_vec_encode_gamma1(const sword32* z, byte l, int bits, if (bits == MLDSA_GAMMA1_BITS_19) { unsigned int e = MLDSA_GAMMA1_19_ENC_BITS / 2 * MLDSA_N / 4; /* Step 2. For each polynomial of vector. */ - for (i = 0; i < l; i++) { - mldsa_encode_gamma1_19_bits(z, s); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_encode_gamma1_19_bits(z, s); /* Move to next place to encode to. */ s += e; /* Next polynomial. */ @@ -2108,6 +2200,7 @@ static void mldsa_vec_encode_gamma1(const sword32* z, byte l, int bits, } } #endif + return ret; } #endif /* WOLFSSL_MLDSA_SIGN_SMALL_MEM */ @@ -2393,10 +2486,14 @@ static void mldsa_decode_gamma1_c(const byte* s, int bits, sword32* z) * @param [in] bits Number of bits used in encoding - GAMMA1 bits. * @param [out] z Polynomial to fill. */ -static void mldsa_decode_gamma1(const byte* s, int bits, sword32* z) +static int mldsa_decode_gamma1(const byte* s, int bits, sword32* z) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; if (bits == MLDSA_GAMMA1_BITS_17) { wc_mldsa_decode_gamma1_17_avx2(s, z); } @@ -2410,6 +2507,7 @@ static void mldsa_decode_gamma1(const byte* s, int bits, sword32* z) { mldsa_decode_gamma1_c(s, bits, z); } + return ret; } #endif @@ -2431,9 +2529,10 @@ static void mldsa_decode_gamma1(const byte* s, int bits, sword32* z) #ifndef WOLFSSL_MLDSA_VERIFY_SMALLEST_MEM /* The smallest-mem verify streams z one polynomial at a time with * mldsa_decode_gamma1() directly, so the whole-vector wrapper is unused. */ -static void mldsa_vec_decode_gamma1(const byte* x, byte l, int bits, +static int mldsa_vec_decode_gamma1(const byte* x, byte l, int bits, sword32* z) { + int ret = 0; unsigned int i; unsigned int e = MLDSA_N / 8 * (unsigned int)(bits + 1); @@ -2441,7 +2540,10 @@ static void mldsa_vec_decode_gamma1(const byte* x, byte l, int bits, i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* Two polynomials per call; an odd trailing one falls through. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i + 1 < l; i += 2) { if (bits == MLDSA_GAMMA1_BITS_17) { wc_mldsa_decode_gamma1_17_x2_avx512(x, x + e, z, z + MLDSA_N); @@ -2456,13 +2558,14 @@ static void mldsa_vec_decode_gamma1(const byte* x, byte l, int bits, } #endif /* Step 3: For each polynomial of vector. */ - for (; i < l; i++) { + for (; (ret == 0) && (i < l); i++) { /* Step 4: Unpack a polynomial. */ - mldsa_decode_gamma1(x, bits, z); + ret = mldsa_decode_gamma1(x, bits, z); /* Move pointers on to next polynomial. */ x += e; z += MLDSA_N; } + return ret; } #endif /* !WOLFSSL_MLDSA_VERIFY_SMALLEST_MEM */ #endif @@ -2534,12 +2637,16 @@ static void mldsa_encode_w1_88_c(const sword32* w1, byte* w1e) * @param [in] w1 Vector of polynomials to encode. * @param [out] w1e Buffer to encode into. */ -static void mldsa_encode_w1_88(const sword32* w1, byte* w1e) +static int mldsa_encode_w1_88(const sword32* w1, byte* w1e) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP /* Tests call this without wc_MlDsaKey_Init. */ cpuid_get_flags_ex(&cpuid_flags); - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_encode_w1_88_avx2(w1, w1e); RESTORE_VECTOR_REGISTERS(); } @@ -2548,11 +2655,12 @@ static void mldsa_encode_w1_88(const sword32* w1, byte* w1e) { mldsa_encode_w1_88_c(w1, w1e); } + return ret; } -WOLFSSL_TEST_VIS void wc_mldsa_encode_w1_88(const sword32* w1, byte* w1e) +WOLFSSL_TEST_VIS int wc_mldsa_encode_w1_88(const sword32* w1, byte* w1e) { - mldsa_encode_w1_88(w1, w1e); + return mldsa_encode_w1_88(w1, w1e); } #endif /* !WOLFSSL_NO_ML_DSA_44 */ @@ -2613,12 +2721,16 @@ static void mldsa_encode_w1_32_c(const sword32* w1, byte* w1e) * @param [in] w1 Vector of polynomials to encode. * @param [out] w1e Buffer to encode into. */ -static void mldsa_encode_w1_32(const sword32* w1, byte* w1e) +static int mldsa_encode_w1_32(const sword32* w1, byte* w1e) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP /* Tests call this without wc_MlDsaKey_Init. */ cpuid_get_flags_ex(&cpuid_flags); - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_encode_w1_32_avx2(w1, w1e); RESTORE_VECTOR_REGISTERS(); } @@ -2627,11 +2739,12 @@ static void mldsa_encode_w1_32(const sword32* w1, byte* w1e) { mldsa_encode_w1_32_c(w1, w1e); } + return ret; } -WOLFSSL_TEST_VIS void wc_mldsa_encode_w1_32(const sword32* w1, byte* w1e) +WOLFSSL_TEST_VIS int wc_mldsa_encode_w1_32(const sword32* w1, byte* w1e) { - mldsa_encode_w1_32(w1, w1e); + return mldsa_encode_w1_32(w1, w1e); } #endif #endif @@ -2654,9 +2767,10 @@ WOLFSSL_TEST_VIS void wc_mldsa_encode_w1_32(const sword32* w1, byte* w1e) * @param [in] gamma2 Maximum value in range. * @param [out] w1e Buffer to encode into. */ -static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, +static int mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, byte* w1e) { + int ret = 0; unsigned int i; (void)k; @@ -2669,8 +2783,10 @@ static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512_VBMI) /* vpermb gathers the encoded bytes of a whole register, so the VBMI * encoder needs no polynomial pairing and does the entire vector. */ - if (IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX512_VBMI(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i < PARAMS_ML_DSA_44_K; i++) { wc_mldsa_encode_w1_88_avx512_vbmi(w1, w1e); w1 += MLDSA_N; @@ -2681,8 +2797,10 @@ static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, #endif #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* Two polynomials per pass; an odd trailing one falls through. */ - if ((i == 0) && USE_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if ((i == 0) && USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i + 1 < PARAMS_ML_DSA_44_K; i += 2) { wc_mldsa_encode_w1_88_x2_avx512(w1, w1 + MLDSA_N, w1e, w1e + e); @@ -2693,8 +2811,8 @@ static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, } #endif /* Step 2. For each polynomial of vector. */ - for (; i < PARAMS_ML_DSA_44_K; i++) { - mldsa_encode_w1_88(w1, w1e); + for (; (ret == 0) && (i < PARAMS_ML_DSA_44_K); i++) { + ret = mldsa_encode_w1_88(w1, w1e); /* Next polynomial. */ w1 += MLDSA_N; w1e += e; @@ -2708,8 +2826,10 @@ static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* One polynomial per pass - vpmovqb needs no lane pairing. */ - if (USE_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i < k; i++) { wc_mldsa_encode_w1_32_avx512(w1, w1e); w1 += MLDSA_N; @@ -2719,8 +2839,8 @@ static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, } #endif /* Step 2. For each polynomial of vector. */ - for (; i < k; i++) { - mldsa_encode_w1_32(w1, w1e); + for (; (ret == 0) && (i < k); i++) { + ret = mldsa_encode_w1_32(w1, w1e); /* Next polynomial. */ w1 += MLDSA_N; w1e += e; @@ -2730,6 +2850,7 @@ static void mldsa_vec_encode_w1(const sword32* w1, byte k, sword32 gamma2, #endif { } + return ret; } #endif @@ -3612,8 +3733,10 @@ static int mldsa_expand_a(wc_Shake* shake128, const byte* pub_seed, #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 /* Eight SHAKE-128 instances per permutation instead of four. Handles any * k x l, so it comes before the fixed-size AVX2 implementations. */ - if (USE_INTEL_AVX512(cpuid_flags) && IS_INTEL_BMI2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags) && IS_INTEL_BMI2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; XMEMCPY(seed, pub_seed, MLDSA_PUB_SEED_SZ); ret = wc_mldsa_gen_matrix_avx512(a, seed, k, l, heap); RESTORE_VECTOR_REGISTERS(); @@ -3622,7 +3745,10 @@ static int mldsa_expand_a(wc_Shake* shake128, const byte* pub_seed, #endif #ifndef WOLFSSL_NO_ML_DSA_44 if ((k == 4) && (l == 4) && IS_INTEL_AVX2(cpuid_flags) && - IS_INTEL_BMI2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_BMI2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; XMEMCPY(seed, pub_seed, MLDSA_PUB_SEED_SZ); ret = wc_mldsa_gen_matrix_4x4_avx2(a, seed); RESTORE_VECTOR_REGISTERS(); @@ -3631,7 +3757,10 @@ static int mldsa_expand_a(wc_Shake* shake128, const byte* pub_seed, #endif #ifndef WOLFSSL_NO_ML_DSA_65 if ((k == 6) && (l == 5) && IS_INTEL_AVX2(cpuid_flags) && - IS_INTEL_BMI2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_BMI2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; XMEMCPY(seed, pub_seed, MLDSA_PUB_SEED_SZ); ret = wc_mldsa_gen_matrix_6x5_avx2(a, seed); RESTORE_VECTOR_REGISTERS(); @@ -3640,7 +3769,10 @@ static int mldsa_expand_a(wc_Shake* shake128, const byte* pub_seed, #endif #ifndef WOLFSSL_NO_ML_DSA_87 if ((k == 8) && (l == 7) && IS_INTEL_AVX2(cpuid_flags) && - IS_INTEL_BMI2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_BMI2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; XMEMCPY(seed, pub_seed, MLDSA_PUB_SEED_SZ); ret = wc_mldsa_gen_matrix_8x7_avx2(a, seed); RESTORE_VECTOR_REGISTERS(); @@ -4700,7 +4832,10 @@ static int mldsa_expand_s(wc_Shake* shake256, byte* priv_seed, byte eta, #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 /* Eight SHAKE-256 instances per permutation instead of four. Handles any * vector dimensions, so it comes before the fixed-size AVX2 code. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = wc_mldsa_gen_s_avx512(s1, s1Len, s2, s2Len, eta, priv_seed, heap); RESTORE_VECTOR_REGISTERS(); @@ -4708,30 +4843,36 @@ static int mldsa_expand_s(wc_Shake* shake256, byte* priv_seed, byte eta, else #endif #ifndef WOLFSSL_NO_ML_DSA_44 - if ((s1Len == 4) && IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) - { + if ((s1Len == 4) && IS_INTEL_AVX2(cpuid_flags)) { sword32* s[2] = { s1, s2 }; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; ret = wc_mldsa_gen_s_4_4_avx2(s, priv_seed); RESTORE_VECTOR_REGISTERS(); } else #endif #ifndef WOLFSSL_NO_ML_DSA_65 - if ((s1Len == 5) && IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) - { + if ((s1Len == 5) && IS_INTEL_AVX2(cpuid_flags)) { sword32* s[2] = { s1, s2 }; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; ret = wc_mldsa_gen_s_5_6_avx2(s, priv_seed); RESTORE_VECTOR_REGISTERS(); } else #endif #ifndef WOLFSSL_NO_ML_DSA_87 - if ((s1Len == 7) && IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) - { + if ((s1Len == 7) && IS_INTEL_AVX2(cpuid_flags)) { sword32* s[2] = { s1, s2 }; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; ret = wc_mldsa_gen_s_7_8_avx2(s, priv_seed); RESTORE_VECTOR_REGISTERS(); } @@ -5143,7 +5284,7 @@ static int mldsa_vec_expand_mask_c(wc_Shake* shake256, byte* seed, MLDSA_MAX_V_BLOCKS); if (ret == 0) { /* Decode v into polynomial. */ - mldsa_decode_gamma1(v, gamma1_bits, y); + ret = mldsa_decode_gamma1(v, gamma1_bits, y); /* Next polynomial. */ y += MLDSA_N; } @@ -5178,15 +5319,19 @@ static int mldsa_vec_expand_mask(wc_Shake* shake256, byte* seed, #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 /* Whole vector in one eight-way run, whatever the dimension. */ - if (USE_INTEL_AVX512(cpuid_flags) && IS_INTEL_BMI2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags) && IS_INTEL_BMI2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = wc_mldsa_gen_y_avx512(y, seed, kappa, gamma1_bits, l, heap); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && IS_INTEL_BMI2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags) && IS_INTEL_BMI2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; #ifndef WOLFSSL_NO_ML_DSA_44 if (l == 4) { ret = wc_mldsa_gen_y_4_avx2(y, seed, kappa); @@ -5348,8 +5493,10 @@ static int mldsa_sample_in_ball_ex(int level, wc_Shake* shake256, if (k == MLDSA_GEN_C_BLOCK_BYTES) { /* Generate a new block. */ #ifndef WC_SHA3_NO_ASM - if (SHA3_USE_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (SHA3_USE_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -5619,12 +5766,16 @@ static void mldsa_vec_decompose_c(const sword32* r, byte k, sword32 gamma2, * @param [out] r0 Low parts in vector of polynomials. * @param [out] r1 High parts in vector of polynomials. */ -static void mldsa_vec_decompose(const sword32* r, byte k, sword32 gamma2, +static int mldsa_vec_decompose(const sword32* r, byte k, sword32 gamma2, sword32* r0, sword32* r1) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; #ifndef WOLFSSL_NO_ML_DSA_44 if (gamma2 == MLDSA_Q_LOW_88) { wc_mldsa_decompose_q88_avx512(r, r0, r1); @@ -5639,7 +5790,10 @@ static void mldsa_vec_decompose(const sword32* r, byte k, sword32 gamma2, } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; #ifndef WOLFSSL_NO_ML_DSA_44 if (gamma2 == MLDSA_Q_LOW_88) { wc_mldsa_decompose_q88_avx2(r, r0, r1); @@ -5657,6 +5811,7 @@ static void mldsa_vec_decompose(const sword32* r, byte k, sword32 gamma2, { mldsa_vec_decompose_c(r, k, gamma2, r0, r1); } + return ret; } #endif @@ -5740,29 +5895,45 @@ static int mldsa_vec_check_low_c(const sword32* a, byte l, sword32 hi) * @param [in] a Vector of polynomials. * @param [in] l Dimension of vector. * @param [in] hi Largest value in range. + * @param [out] valid 1 when every value is in range, 0 otherwise. + * @return 0 on success, or the error from a refused vector-register save. */ -static int mldsa_vec_check_low(const sword32* a, byte l, sword32 hi) +static int mldsa_vec_check_low(const sword32* a, byte l, sword32 hi, + int* valid) { - int ret; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = wc_mldsa_vec_check_low_avx512(a, l, hi); + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) { + /* Say "not valid" so a refused save cannot read as a pass. */ + *valid = 0; + return svr_ret; + } + *valid = wc_mldsa_vec_check_low_avx512(a, l, hi); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = wc_mldsa_vec_check_low_avx2(a, l, hi); + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) { + /* Say "not valid" so a refused save cannot read as a pass. */ + *valid = 0; + return svr_ret; + } + *valid = wc_mldsa_vec_check_low_avx2(a, l, hi); RESTORE_VECTOR_REGISTERS(); } else #endif { - ret = mldsa_vec_check_low_c(a, l, hi); + *valid = mldsa_vec_check_low_c(a, l, hi); } - return ret; + return 0; } #endif @@ -5808,21 +5979,28 @@ static int mldsa_vec_check_low(const sword32* a, byte l, sword32 hi) * @param [in] w1 Vector of polynomials that is high part of w. * @param [out] h Encoded hints. * @param [in, out] idxp Index to write next hint into. - * return Number of hints on success. - * return Falsam of -1 when too many hints. + * @param [out] valid 1 when the hints fit, 0 when there are too many. + * @return 0 on success, or the error from a refused vector-register save. */ static int mldsa_make_hint_88(const sword32* s, const sword32* w1, byte* h, - byte *idxp) + byte *idxp, int* valid) { unsigned int j; byte idx; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_mldsa_make_hint_88_avx512(s, w1, PARAMS_ML_DSA_44_OMEGA, - h, idxp); + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) { + /* Say "not valid" so a refused save cannot read as a pass. */ + *valid = 0; + return svr_ret; + } + *valid = (wc_mldsa_make_hint_88_avx512(s, w1, PARAMS_ML_DSA_44_OMEGA, + h, idxp) == 0); RESTORE_VECTOR_REGISTERS(); - return ret; + return 0; } #endif @@ -5845,12 +6023,14 @@ static int mldsa_make_hint_88(const sword32* s, const sword32* w1, byte* h, /* Alg 2, Step 27: If there are too many hints, return * falsam of -1. */ if (idx > PARAMS_ML_DSA_44_OMEGA) { - return -1; + *valid = 0; + return 0; } } } *idxp = idx; + *valid = 1; return 0; } #endif @@ -5891,11 +6071,11 @@ static int mldsa_make_hint_88(const sword32* s, const sword32* w1, byte* h, * @param [in] omega Maximum number of hints allowed. * @param [out] h Encoded hints. * @param [in, out] idxp Index to write next hint into. - * return Number of hints on success. - * return Falsam of -1 when too many hints. + * @param [out] valid 1 when the hints fit, 0 when there are too many. + * @return 0 on success, or the error from a refused vector-register save. */ static int mldsa_make_hint_32(const sword32* s, const sword32* w1, - byte omega, byte* h, byte *idxp) + byte omega, byte* h, byte *idxp, int* valid) { unsigned int j; byte idx; @@ -5903,10 +6083,17 @@ static int mldsa_make_hint_32(const sword32* s, const sword32* w1, (void)omega; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_mldsa_make_hint_32_avx512(s, w1, omega, h, idxp); + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) { + /* Say "not valid" so a refused save cannot read as a pass. */ + *valid = 0; + return svr_ret; + } + *valid = (wc_mldsa_make_hint_32_avx512(s, w1, omega, h, idxp) == 0); RESTORE_VECTOR_REGISTERS(); - return ret; + return 0; } #endif @@ -5929,12 +6116,14 @@ static int mldsa_make_hint_32(const sword32* s, const sword32* w1, /* Alg 2, Step 27: If there are too many hints, return * falsam of -1. */ if (idx > omega) { - return -1; + *valid = 0; + return 0; } } } *idxp = idx; + *valid = 1; return 0; } #endif @@ -5981,24 +6170,27 @@ static int mldsa_make_hint_32(const sword32* s, const sword32* w1, * @param [in] gamma2 Low-order rounding range, GAMMA2. * @param [in] omega Maximum number of hints allowed. * @param [out] h Encoded hints. - * return Number of hints on success. - * return Falsam of -1 when too many hints. + * @param [out] valid 1 when the hints fit, 0 when there are too many. + * @return 0 on success, or the error from a refused vector-register save. */ static int mldsa_make_hint(const sword32* s, const sword32* w1, byte k, - sword32 gamma2, byte omega, byte* h) + sword32 gamma2, byte omega, byte* h, int* valid) { + int ret = 0; unsigned int i; byte idx = 0; (void)k; (void)omega; + *valid = 1; #ifndef WOLFSSL_NO_ML_DSA_44 if (gamma2 == MLDSA_Q_LOW_88) { /* Alg 14, Step 2: For each polynomial of vector. */ for (i = 0; i < PARAMS_ML_DSA_44_K; i++) { - if (mldsa_make_hint_88(s, w1, h, &idx) == -1) { - return -1; + ret = mldsa_make_hint_88(s, w1, h, &idx, valid); + if ((ret != 0) || (!*valid)) { + return ret; } /* Alg 14, Step 10: Store count of hints for polynomial at end of * list. */ @@ -6014,8 +6206,9 @@ static int mldsa_make_hint(const sword32* s, const sword32* w1, byte k, if (gamma2 == MLDSA_Q_LOW_32) { /* Alg 14, Step 2: For each polynomial of vector. */ for (i = 0; i < k; i++) { - if (mldsa_make_hint_32(s, w1, omega, h, &idx) == -1) { - return -1; + ret = mldsa_make_hint_32(s, w1, omega, h, &idx, valid); + if ((ret != 0) || (!*valid)) { + return ret; } /* Alg 14, Step 10: Store count of hints for polynomial at end of * list. */ @@ -6032,7 +6225,7 @@ static int mldsa_make_hint(const sword32* s, const sword32* w1, byte k, /* Set remaining hints to zero. */ XMEMSET(h + idx, 0, (size_t)(omega - idx)); - return idx; + return 0; } #endif /* !WOLFSSL_MLDSA_SIGN_SMALL_MEM */ @@ -6255,9 +6448,10 @@ static void mldsa_use_hint_32(sword32* w1, const byte* h, byte omega, * @param [in] omega Max number of hints. Hint counts after this index. * @param [in] h Hints to apply. In signature encoding. */ -static void mldsa_vec_use_hint(sword32* w1, byte k, sword32 gamma2, +static int mldsa_vec_use_hint(sword32* w1, byte k, sword32 gamma2, byte omega, const byte* h) { + int ret = 0; unsigned int i; byte o = 0; @@ -6268,14 +6462,19 @@ static void mldsa_vec_use_hint(sword32* w1, byte k, sword32 gamma2, if (gamma2 == MLDSA_Q_LOW_88) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_use_hint_88_avx512(w1, h); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_use_hint_88_avx2(w1, h); RESTORE_VECTOR_REGISTERS(); } @@ -6294,14 +6493,19 @@ static void mldsa_vec_use_hint(sword32* w1, byte k, sword32 gamma2, if (gamma2 == MLDSA_Q_LOW_32) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_use_hint_32_avx512(w1, k, h); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_use_hint_32_avx2(w1, k, h); RESTORE_VECTOR_REGISTERS(); } @@ -6316,6 +6520,7 @@ static void mldsa_vec_use_hint(sword32* w1, byte k, sword32 gamma2, } } #endif + return ret; } #endif #endif /* !WOLFSSL_MLDSA_NO_VERIFY */ @@ -6862,10 +7067,14 @@ static void mldsa_ntt_c(sword32* r) * * @param [in, out] r Polynomial to transform. */ -static void mldsa_ntt(sword32* r) +static int mldsa_ntt(sword32* r) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; MLDSA_NTT_AVX512(r); RESTORE_VECTOR_REGISTERS(); } @@ -6873,7 +7082,10 @@ static void mldsa_ntt(sword32* r) #endif #ifdef USE_INTEL_SPEEDUP /* MLDSA_NTT_AVX2: see the flavor-selection note by its definition. */ - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; MLDSA_NTT_AVX2(r); RESTORE_VECTOR_REGISTERS(); } @@ -6882,6 +7094,7 @@ static void mldsa_ntt(sword32* r) { mldsa_ntt_c(r); } + return ret; } #endif @@ -6896,17 +7109,24 @@ static void mldsa_ntt(sword32* r) * * @param [in, out] r Polynomial to transform. */ -static void mldsa_ntt_full(sword32* r) +static int mldsa_ntt_full(sword32* r) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_ntt_full_1p_avx512(r); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_ntt_full_avx2(r); RESTORE_VECTOR_REGISTERS(); } @@ -6915,6 +7135,7 @@ static void mldsa_ntt_full(sword32* r) { mldsa_ntt_c(r); } + return ret; } #endif @@ -6927,14 +7148,18 @@ static void mldsa_ntt_full(sword32* r) * @param [in, out] r Vector of polynomials to transform. * @param [in] l Dimension of polynomial. */ -static void mldsa_vec_ntt(sword32* r, byte l) +static int mldsa_vec_ntt(sword32* r, byte l) { + int ret = 0; unsigned int i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* One polynomial per call, so there is no odd trailing polynomial to * hand back to mldsa_ntt() below. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i < l; i++) { /* MLDSA_NTT_AVX512: see the flavor-selection note by its * definition. */ @@ -6944,10 +7169,11 @@ static void mldsa_vec_ntt(sword32* r, byte l) RESTORE_VECTOR_REGISTERS(); } #endif - for (; i < l; i++) { - mldsa_ntt(r); + for (; (ret == 0) && (i < l); i++) { + ret = mldsa_ntt(r); r += MLDSA_N; } + return ret; } #endif #endif @@ -6964,14 +7190,16 @@ static void mldsa_vec_ntt(sword32* r, byte l) * @param [in, out] r Vector of polynomials to transform. * @param [in] l Dimension of polynomial. */ -static void mldsa_vec_ntt_full(sword32* r, byte l) +static int mldsa_vec_ntt_full(sword32* r, byte l) { + int ret = 0; unsigned int i; - for (i = 0; i < l; i++) { - mldsa_ntt_full(r); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_ntt_full(r); r += MLDSA_N; } + return ret; } #endif @@ -7342,10 +7570,14 @@ static void mldsa_ntt_small_c(sword32* r) * * @param [in, out] r Polynomial to transform, coefficients in -26..26. */ -static void mldsa_ntt_small(sword32* r) +static int mldsa_ntt_small(sword32* r) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; MLDSA_NTT_SMALL_AVX512(r); RESTORE_VECTOR_REGISTERS(); } @@ -7354,7 +7586,10 @@ static void mldsa_ntt_small(sword32* r) #ifdef USE_INTEL_SPEEDUP /* MLDSA_NTT_SMALL_AVX2: see the flavor-selection note by its * definition. */ - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; MLDSA_NTT_SMALL_AVX2(r); RESTORE_VECTOR_REGISTERS(); } @@ -7363,6 +7598,7 @@ static void mldsa_ntt_small(sword32* r) { mldsa_ntt_small_c(r); } + return ret; } #endif @@ -7379,17 +7615,24 @@ static void mldsa_ntt_small(sword32* r) * * @param [in, out] r Polynomial to transform, coefficients in -26..26. */ -static void mldsa_ntt_small_full(sword32* r) +static int mldsa_ntt_small_full(sword32* r) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_ntt_small_full_1p_avx512(r); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_ntt_small_full_avx2(r); RESTORE_VECTOR_REGISTERS(); } @@ -7398,6 +7641,7 @@ static void mldsa_ntt_small_full(sword32* r) { mldsa_ntt_small_c(r); } + return ret; } #endif @@ -7410,14 +7654,16 @@ static void mldsa_ntt_small_full(sword32* r) * @param [in, out] r Vector of polynomials to transform. * @param [in] l Dimension of polynomial. */ -static void mldsa_vec_ntt_small(sword32* r, byte l) +static int mldsa_vec_ntt_small(sword32* r, byte l) { + int ret = 0; unsigned int i; - for (i = 0; i < l; i++) { - mldsa_ntt_small(r); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_ntt_small(r); r += MLDSA_N; } + return ret; } #endif @@ -7428,14 +7674,16 @@ static void mldsa_vec_ntt_small(sword32* r, byte l) * @param [in, out] r Vector of polynomials to transform. * @param [in] l Dimension of polynomial. */ -static void mldsa_vec_ntt_small_full(sword32* r, byte l) +static int mldsa_vec_ntt_small_full(sword32* r, byte l) { + int ret = 0; unsigned int i; - for (i = 0; i < l; i++) { - mldsa_ntt_small_full(r); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_ntt_small_full(r); r += MLDSA_N; } + return ret; } #endif @@ -7891,10 +8139,14 @@ static void mldsa_invntt_c(sword32* r) * * @param [in, out] r Polynomial to transform. */ -static void mldsa_invntt(sword32* r) +static int mldsa_invntt(sword32* r) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; /* MLDSA_INVNTT_AVX512: see the flavor-selection note by its * definition. */ MLDSA_INVNTT_AVX512(r); @@ -7904,7 +8156,10 @@ static void mldsa_invntt(sword32* r) #endif #ifdef USE_INTEL_SPEEDUP /* MLDSA_INVNTT_AVX2: see the flavor-selection note by its definition. */ - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; MLDSA_INVNTT_AVX2(r); RESTORE_VECTOR_REGISTERS(); } @@ -7913,6 +8168,7 @@ static void mldsa_invntt(sword32* r) { mldsa_invntt_c(r); } + return ret; } #endif @@ -7920,17 +8176,24 @@ static void mldsa_invntt(sword32* r) * * @param [in, out] r Polynomial to transform. */ -static void mldsa_invntt_full(sword32* r) +static int mldsa_invntt_full(sword32* r) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_invntt_full_1p_avx512(r); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef USE_INTEL_SPEEDUP - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_invntt_full_avx2(r); RESTORE_VECTOR_REGISTERS(); } @@ -7939,6 +8202,7 @@ static void mldsa_invntt_full(sword32* r) { mldsa_invntt_c(r); } + return ret; } #if !defined(WOLFSSL_MLDSA_NO_MAKE_KEY) || \ @@ -7952,14 +8216,18 @@ static void mldsa_invntt_full(sword32* r) * @param [in, out] r Vector of polynomials to transform. * @param [in] l Dimension of polynomial. */ -static void mldsa_vec_invntt_full(sword32* r, byte l) +static int mldsa_vec_invntt_full(sword32* r, byte l) { + int ret = 0; unsigned int i = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* One polynomial per call, so there is no odd trailing polynomial to * hand back to mldsa_invntt_full() below. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (; i < l; i++) { wc_mldsa_invntt_full_1p_avx512(r); r += MLDSA_N; @@ -7967,10 +8235,11 @@ static void mldsa_vec_invntt_full(sword32* r, byte l) RESTORE_VECTOR_REGISTERS(); } #endif - for (; i < l; i++) { - mldsa_invntt_full(r); + for (; (ret == 0) && (i < l); i++) { + ret = mldsa_invntt_full(r); r += MLDSA_N; } + return ret; } #endif @@ -8150,13 +8419,18 @@ static void mldsa_matrix_mul_c(sword32* r, const sword32* m, * @param [in] k First dimension of matrix and dimension of result. * @param [in] l Second dimension of matrix and dimension of v. */ -static void mldsa_matrix_mul(sword32* r, const sword32* m, const sword32* v, +static int mldsa_matrix_mul(sword32* r, const sword32* m, const sword32* v, byte k, byte l) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; if (l == 4) { for (i = 0; i < k; i++) { wc_mldsa_mul_vec_4_avx512(r, m, v); @@ -8182,8 +8456,12 @@ static void mldsa_matrix_mul(sword32* r, const sword32* m, const sword32* v, } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) + return svr_ret; if (l == 4) { for (i = 0; i < k; i++) { wc_mldsa_mul_vec_4_avx2(r, m, v); @@ -8212,6 +8490,7 @@ static void mldsa_matrix_mul(sword32* r, const sword32* m, const sword32* v, { mldsa_matrix_mul_c(r, m, v, k, l); } + return ret; } #endif @@ -8271,17 +8550,24 @@ static void mldsa_mul_c(sword32* r, sword32* a, sword32* b) * @param [in] a Polynomial * @param [in] b Polynomial. */ -static void mldsa_mul(sword32* r, sword32* a, sword32* b) +static int mldsa_mul(sword32* r, sword32* a, sword32* b) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_mul_avx512(r, a, b); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_mul_avx2(r, a, b); RESTORE_VECTOR_REGISTERS(); } @@ -8290,6 +8576,7 @@ static void mldsa_mul(sword32* r, sword32* a, sword32* b) { mldsa_mul_c(r, a, b); } + return ret; } #ifndef WOLFSSL_MLDSA_SIGN_SMALL_MEM @@ -8300,11 +8587,15 @@ static void mldsa_mul(sword32* r, sword32* a, sword32* b) * @param [in] c Challenge polynomial in NTT form. * @param [in] v Polynomial of vector in NTT form. */ -static void mldsa_mul_invntt(sword32* r, sword32* c, sword32* v) +static int mldsa_mul_invntt(sword32* r, sword32* c, sword32* v) { + int ret = 0; #if defined(USE_INTEL_SPEEDUP) && defined(WOLFSSL_MLDSA_HAVE_INTEL_AVX512) /* Both steps under one save/restore of the vector registers. */ - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_mul_avx512(r, c, v); /* MLDSA_INVNTT_AVX512: see the flavor-selection note by its * definition. wc_mldsa_mul_avx512() is positionwise, so its output @@ -8315,9 +8606,14 @@ static void mldsa_mul_invntt(sword32* r, sword32* c, sword32* v) else #endif { - mldsa_mul(r, c, v); - mldsa_invntt(r); + if (ret == 0) { + ret = mldsa_mul(r, c, v); + } + if (ret == 0) { + ret = mldsa_invntt(r); + } } + return ret; } #endif /* !WOLFSSL_MLDSA_SIGN_SMALL_MEM */ #endif @@ -8331,13 +8627,17 @@ static void mldsa_mul_invntt(sword32* r, sword32* c, sword32* v) * @param [in] b Vector of polynomials. * @param [in] l Dimension of vectors. */ -static void mldsa_vec_mul(sword32* r, sword32* a, sword32* b, byte l) +static int mldsa_vec_mul(sword32* r, sword32* a, sword32* b, byte l) { + int ret = 0; byte i; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < l; i++) { wc_mldsa_mul_avx512(r, a, b); r += MLDSA_N; @@ -8347,7 +8647,10 @@ static void mldsa_vec_mul(sword32* r, sword32* a, sword32* b, byte l) } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < l; i++) { wc_mldsa_mul_avx2(r, a, b); r += MLDSA_N; @@ -8364,6 +8667,7 @@ static void mldsa_vec_mul(sword32* r, sword32* a, sword32* b, byte l) b += MLDSA_N; } } + return ret; } #endif #endif @@ -8403,17 +8707,24 @@ static void mldsa_poly_red_c(sword32* a) * * @param [in, out] a Polynomial. */ -static void mldsa_poly_red(sword32* a) +static int mldsa_poly_red(sword32* a) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_red_avx512(a); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_red_avx2(a); RESTORE_VECTOR_REGISTERS(); } @@ -8422,6 +8733,7 @@ static void mldsa_poly_red(sword32* a) { mldsa_poly_red_c(a); } + return ret; } #if (defined(WOLFSSL_MLDSA_SMALL) && \ @@ -8436,14 +8748,16 @@ static void mldsa_poly_red(sword32* a) * @param [in, out] a Vector of polynomials. * @param [in] l Dimension of vector. */ -static void mldsa_vec_red(sword32* a, byte l) +static int mldsa_vec_red(sword32* a, byte l) { + int ret = 0; byte i; - for (i = 0; i < l; i++) { - mldsa_poly_red(a); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_poly_red(a); a += MLDSA_N; } + return ret; } #endif #endif @@ -8483,17 +8797,24 @@ static void mldsa_sub_c(sword32* r, const sword32* a) * @param [out] r Polynomial to subtract from. * @param [in] a Polynomial to subtract. */ -static void mldsa_sub(sword32* r, const sword32* a) +static int mldsa_sub(sword32* r, const sword32* a) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_sub_avx512(r, a); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_sub_avx2(r, a); RESTORE_VECTOR_REGISTERS(); } @@ -8502,6 +8823,7 @@ static void mldsa_sub(sword32* r, const sword32* a) { mldsa_sub_c(r, a); } + return ret; } #if defined(WOLFSSL_MLDSA_CHECK_KEY) || \ @@ -8513,15 +8835,17 @@ static void mldsa_sub(sword32* r, const sword32* a) * @param [in] a Vector of polynomials to subtract. * @param [in] l Dimension of vectors. */ -static void mldsa_vec_sub(sword32* r, const sword32* a, byte l) +static int mldsa_vec_sub(sword32* r, const sword32* a, byte l) { + int ret = 0; byte i; - for (i = 0; i < l; i++) { - mldsa_sub(r, a); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_sub(r, a); r += MLDSA_N; a += MLDSA_N; } + return ret; } #endif #endif @@ -8558,17 +8882,24 @@ static void mldsa_add_c(sword32* r, const sword32* a) * @param [out] r Polynomial to add to. * @param [in] a Polynomial to add. */ -static void mldsa_add(sword32* r, const sword32* a) +static int mldsa_add(sword32* r, const sword32* a) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_add_avx512(r, a); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_add_avx2(r, a); RESTORE_VECTOR_REGISTERS(); } @@ -8577,6 +8908,7 @@ static void mldsa_add(sword32* r, const sword32* a) { mldsa_add_c(r, a); } + return ret; } #if !defined(WOLFSSL_MLDSA_NO_MAKE_KEY) || \ @@ -8589,15 +8921,17 @@ static void mldsa_add(sword32* r, const sword32* a) * @param [in] a Vector of polynomials to add. * @param [in] l Dimension of vectors. */ -static void mldsa_vec_add(sword32* r, const sword32* a, byte l) +static int mldsa_vec_add(sword32* r, const sword32* a, byte l) { + int ret = 0; byte i; - for (i = 0; i < l; i++) { - mldsa_add(r, a); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_add(r, a); r += MLDSA_N; a += MLDSA_N; } + return ret; } #endif @@ -8636,17 +8970,24 @@ static void mldsa_make_pos_c(sword32* a) * * @param [in, out] a Polynomial. */ -static void mldsa_make_pos(sword32* a) +static int mldsa_make_pos(sword32* a) { + int ret = 0; #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLDSA_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_make_pos_avx512(a); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; wc_mldsa_poly_make_pos_avx2(a); RESTORE_VECTOR_REGISTERS(); } @@ -8655,6 +8996,7 @@ static void mldsa_make_pos(sword32* a) { mldsa_make_pos_c(a); } + return ret; } #if !defined(WOLFSSL_MLDSA_NO_MAKE_KEY) || \ @@ -8666,14 +9008,16 @@ static void mldsa_make_pos(sword32* a) * @param [in, out] a Vector of polynomials. * @param [in] l Dimension of vector. */ -static void mldsa_vec_make_pos(sword32* a, byte l) +static int mldsa_vec_make_pos(sword32* a, byte l) { + int ret = 0; byte i; - for (i = 0; i < l; i++) { - mldsa_make_pos(a); + for (i = 0; (ret == 0) && (i < l); i++) { + ret = mldsa_make_pos(a); a += MLDSA_N; } + return ret; } #endif @@ -8854,28 +9198,48 @@ static int mldsa_make_key_from_seed(wc_MlDsaKey* key, const byte* seed) /* Step 9: Move k down to after public seed. */ XMEMCPY(k, k + MLDSA_PRIV_SEED_SZ, MLDSA_K_SZ); /* Step 9. Alg 24 Steps 2-4: Encode s1 into private key. */ - mldsa_vec_encode_eta_bits(s1, params->l, params->eta, s1p); + if (ret == 0) { + ret = mldsa_vec_encode_eta_bits(s1, params->l, params->eta, s1p); + } /* Step 9. Alg 24 Steps 5-7: Encode s2 into private key. */ - mldsa_vec_encode_eta_bits(s2, params->k, params->eta, s2p); + if (ret == 0) { + ret = mldsa_vec_encode_eta_bits(s2, params->k, params->eta, s2p); + } /* Step 5: t <- NTT-1(A_circum o NTT(s1)) + s2 */ - mldsa_vec_ntt_small_full(s1, params->l); - mldsa_matrix_mul(t, a, s1, params->k, params->l); + if (ret == 0) { + ret = mldsa_vec_ntt_small_full(s1, params->l); + } + if (ret == 0) { + ret = mldsa_matrix_mul(t, a, s1, params->k, params->l); + } #ifdef WOLFSSL_MLDSA_SMALL - mldsa_vec_red(t, params->k); + if (ret == 0) { + ret = mldsa_vec_red(t, params->k); + } #endif - mldsa_vec_invntt_full(t, params->k); - mldsa_vec_add(t, s2, params->k); + if (ret == 0) { + ret = mldsa_vec_invntt_full(t, params->k); + } + if (ret == 0) { + ret = mldsa_vec_add(t, s2, params->k); + } /* Make positive for decomposing. */ - mldsa_vec_make_pos(t, params->k); + if (ret == 0) { + ret = mldsa_vec_make_pos(t, params->k); + } /* Step 6, Step 7, Step 9. Alg 22 Steps 2-4, Alg 24 Steps 8-10. * Decompose t in t0 and t1 and encode into public and private key. */ - mldsa_vec_encode_t0_t1(t, params->k, t0, t1); - /* Step 8. Alg 24, Step 1: Hash public key into private key. */ - ret = mldsa_shake256(&key->shake, key->p, params->pkSz, tr, - MLDSA_TR_SZ); + if (ret == 0) { + ret = mldsa_vec_encode_t0_t1(t, params->k, t0, t1); + } + if (ret == 0) { + /* Step 8. Alg 24, Step 1: Hash public key into private key. */ + ret = mldsa_shake256(&key->shake, key->p, params->pkSz, tr, + MLDSA_TR_SZ); + } } if (ret == 0) { /* Public key and private key are available. */ @@ -9003,12 +9367,18 @@ static int mldsa_make_key_from_seed(wc_MlDsaKey* key, const byte* seed) /* Step 9: Move k down to after public seed. */ XMEMCPY(k, k + MLDSA_PRIV_SEED_SZ, MLDSA_K_SZ); /* Step 9. Alg 24 Steps 2-4: Encode s1 into private key. */ - mldsa_vec_encode_eta_bits(s1, params->l, params->eta, s1p); + if (ret == 0) { + ret = mldsa_vec_encode_eta_bits(s1, params->l, params->eta, s1p); + } /* Step 9. Alg 24 Steps 5-7: Encode s2 into private key. */ - mldsa_vec_encode_eta_bits(s2, params->k, params->eta, s2p); + if (ret == 0) { + ret = mldsa_vec_encode_eta_bits(s2, params->k, params->eta, s2p); + } /* Step 5: NTT(s1) */ - mldsa_vec_ntt_small_full(s1, params->l); + if (ret == 0) { + ret = mldsa_vec_ntt_small_full(s1, params->l); + } /* Step 5: t <- NTT-1(A_circum o NTT(s1)) + s2 */ XMEMCPY(aseed, pub_seed, MLDSA_PUB_SEED_SZ); for (r = 0; (ret == 0) && (r < params->k); r++) { @@ -9109,10 +9479,16 @@ static int mldsa_make_key_from_seed(wc_MlDsaKey* key, const byte* seed) tt[e] = mldsa_mont_red(t64[e]); } #endif - mldsa_invntt_full(tt); - mldsa_add(tt, s2t); + if (ret == 0) { + ret = mldsa_invntt_full(tt); + } + if (ret == 0) { + ret = mldsa_add(tt, s2t); + } /* Make positive for decomposing. */ - mldsa_make_pos(tt); + if (ret == 0) { + ret = mldsa_make_pos(tt); + } tt += MLDSA_N; s2t += MLDSA_N; @@ -9121,10 +9497,14 @@ static int mldsa_make_key_from_seed(wc_MlDsaKey* key, const byte* seed) /* Step 6, Step 7, Step 9. Alg 22 Steps 2-4, Alg 24 Steps 8-10. * Decompose t in t0 and t1 and encode into public and private key. */ - mldsa_vec_encode_t0_t1(t, params->k, t0, t1); - /* Step 8. Alg 24, Step 1: Hash public key into private key. */ - ret = mldsa_shake256(&key->shake, key->p, params->pkSz, tr, - MLDSA_TR_SZ); + if (ret == 0) { + ret = mldsa_vec_encode_t0_t1(t, params->k, t0, t1); + } + if (ret == 0) { + /* Step 8. Alg 24, Step 1: Hash public key into private key. */ + ret = mldsa_shake256(&key->shake, key->p, params->pkSz, tr, + MLDSA_TR_SZ); + } } if (ret == 0) { /* Public key and private key are available. */ @@ -9265,9 +9645,10 @@ static int mldsa_make_key(wc_MlDsaKey* key, WC_RNG* rng) * @param [out] s2 Vector of polynomials s2. * @param [out] t0 Vector of polynomials t0. */ -static void mldsa_make_priv_vecs(wc_MlDsaKey* key, sword32* s1, +static int mldsa_make_priv_vecs(wc_MlDsaKey* key, sword32* s1, sword32* s2, sword32* t0) { + int ret = 0; const wc_MlDsaParams* params = key->params; const byte* pubSeed = key->k; const byte* k = pubSeed + MLDSA_PUB_SEED_SZ; @@ -9277,21 +9658,36 @@ static void mldsa_make_priv_vecs(wc_MlDsaKey* key, sword32* s1, const byte* t0p = s2p + params->s2EncSz; /* Step 1: Decode s1, s2, t0. */ - mldsa_vec_decode_eta_bits(s1p, params->eta, s1, params->l); - mldsa_vec_decode_eta_bits(s2p, params->eta, s2, params->k); - mldsa_vec_decode_t0(t0p, params->k, t0); + if (ret == 0) { + ret = mldsa_vec_decode_eta_bits(s1p, params->eta, s1, params->l); + } + if (ret == 0) { + ret = mldsa_vec_decode_eta_bits(s2p, params->eta, s2, params->k); + } + if (ret == 0) { + ret = mldsa_vec_decode_t0(t0p, params->k, t0); + } /* Step 2: NTT s1. */ - mldsa_vec_ntt_small(s1, params->l); + if (ret == 0) { + ret = mldsa_vec_ntt_small(s1, params->l); + } /* Step 3: NTT s2. */ - mldsa_vec_ntt_small(s2, params->k); + if (ret == 0) { + ret = mldsa_vec_ntt_small(s2, params->k); + } /* Step 4: NTT t0. */ - mldsa_vec_ntt(t0, params->k); + if (ret == 0) { + ret = mldsa_vec_ntt(t0, params->k); + } #ifdef WC_MLDSA_CACHE_PRIV_VECTORS - /* Private key vectors have been created. */ - key->privVecsSet = 1; + if (ret == 0) { + /* Private key vectors have been created. */ + key->privVecsSet = 1; + } #endif + return ret; } #endif @@ -9481,12 +9877,14 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, #endif { /* Steps 1-4: Decode and NTT vectors s1, s2, and t0. */ - mldsa_make_priv_vecs(key, s1, s2, t0); + ret = mldsa_make_priv_vecs(key, s1, s2, t0); } #ifdef WC_MLDSA_CACHE_MATRIX_A /* Check that we haven't already cached the matrix A. */ - if (!key->aSet) + if ((ret == 0) && (!key->aSet)) +#else + if (ret == 0) #endif { /* Step 5: Create the matrix A from the public seed. */ @@ -9516,12 +9914,15 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, byte* commit = sig; /* Step 12: Compute vector y from private random seed and kappa. */ - mldsa_vec_expand_mask(&key->shake, priv_rand_seed, kappa, + ret = mldsa_vec_expand_mask(&key->shake, priv_rand_seed, kappa, params->gamma1_bits, y, params->l, key->heap); #ifdef WOLFSSL_MLDSA_SIGN_CHECK_Y - valid = mldsa_vec_check_low(y, params->l, - ((sword32)1 << params->gamma1_bits) - params->beta); - if (valid) + if (ret == 0) { + ret = mldsa_vec_check_low(y, params->l, + ((sword32)1 << params->gamma1_bits) - params->beta, + &valid); + } + if ((ret == 0) && valid) #endif { /* Step 13: NTT-1(A o NTT(y)) */ @@ -9534,31 +9935,51 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, } if (ret == 0) { #endif - mldsa_vec_ntt_full(y_ntt, params->l); - mldsa_matrix_mul(w, a, y_ntt, params->k, params->l); + if (ret == 0) { + ret = mldsa_vec_ntt_full(y_ntt, params->l); + } + if (ret == 0) { + ret = mldsa_matrix_mul(w, a, y_ntt, params->k, params->l); + } #ifdef WOLFSSL_MLDSA_SMALL - mldsa_vec_red(w, params->k); + if (ret == 0) { + ret = mldsa_vec_red(w, params->k); + } #endif - mldsa_vec_invntt_full(w, params->k); + if (ret == 0) { + ret = mldsa_vec_invntt_full(w, params->k); + } /* Step 14, Step 22: Make values positive and decompose. */ - mldsa_vec_make_pos(w, params->k); - mldsa_vec_decompose(w, params->k, params->gamma2, w0, w1); + if (ret == 0) { + ret = mldsa_vec_make_pos(w, params->k); + } + if (ret == 0) { + ret = mldsa_vec_decompose(w, params->k, params->gamma2, w0, + w1); + } #ifdef WOLFSSL_MLDSA_SIGN_CHECK_W0 - valid = mldsa_vec_check_low(w0, params->k, - params->gamma2 - params->beta); + if (ret == 0) { + ret = mldsa_vec_check_low(w0, params->k, + params->gamma2 - params->beta, &valid); + } } - if (valid) { + if ((ret == 0) && valid) { #endif /* Step 15: Encode w1. */ WC_ALLOC_VAR_EX(w1e, byte, MLDSA_MAX_W1_ENC_SZ, key->heap, DYNAMIC_TYPE_MLDSA, ret=MEMORY_E); if (WC_VAR_OK(w1e)) { - mldsa_vec_encode_w1(w1, params->k, params->gamma2, w1e); - /* Step 15: Hash mu and encoded w1. - * Step 32: Hash is stored in signature. */ - ret = mldsa_hash256(&key->shake, mu, MLDSA_MU_SZ, - w1e, params->w1EncSz, commit, params->lambda / 4); + if (ret == 0) { + ret = mldsa_vec_encode_w1(w1, params->k, + params->gamma2, w1e); + } + if (ret == 0) { + /* Step 15: Hash mu and encoded w1. + * Step 32: Hash is stored in signature. */ + ret = mldsa_hash256(&key->shake, mu, MLDSA_MU_SZ, + w1e, params->w1EncSz, commit, params->lambda / 4); + } } if (ret == 0) { /* Step 17: Compute c from first 256 bits of commit. */ @@ -9571,46 +9992,66 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, valid = 1; /* Step 18: NTT(c). */ - mldsa_ntt_small(c); + ret = mldsa_ntt_small(c); hi = params->gamma2 - params->beta; - for (i = 0; valid && i < params->k; i++) { + for (i = 0; (ret == 0) && valid && (i < params->k); i++) { /* Step 20: cs2 = NTT-1(c o s2) */ - mldsa_mul_invntt(cs2 + i * MLDSA_N, c, + ret = mldsa_mul_invntt(cs2 + i * MLDSA_N, c, s2 + i * MLDSA_N); /* Step 22: w0 - cs2 */ - mldsa_sub(w0 + i * MLDSA_N, cs2 + i * MLDSA_N); + if (ret == 0) { + ret = mldsa_sub(w0 + i * MLDSA_N, + cs2 + i * MLDSA_N); + } /* Step 23: Check w0 - cs2 has low enough values. */ - valid = mldsa_vec_check_low(w0 + i * MLDSA_N, 1, hi); + if (ret == 0) { + ret = mldsa_vec_check_low(w0 + i * MLDSA_N, 1, hi, + &valid); + } } hi = ((sword32)1 << params->gamma1_bits) - params->beta; - for (i = 0; valid && i < params->l; i++) { + for (i = 0; (ret == 0) && valid && (i < params->l); i++) { /* Step 19: cs1 = NTT-1(c o s1) */ - mldsa_mul_invntt(z + i * MLDSA_N, c, + ret = mldsa_mul_invntt(z + i * MLDSA_N, c, s1 + i * MLDSA_N); /* Step 21: z = y + cs1 */ - mldsa_add(z + i * MLDSA_N, y + i * MLDSA_N); - mldsa_poly_red(z + i * MLDSA_N); + if (ret == 0) { + ret = mldsa_add(z + i * MLDSA_N, y + i * MLDSA_N); + } + if (ret == 0) { + ret = mldsa_poly_red(z + i * MLDSA_N); + } /* Step 23: Check z has low enough values. */ - valid = mldsa_vec_check_low(z + i * MLDSA_N, 1, hi); + if (ret == 0) { + ret = mldsa_vec_check_low(z + i * MLDSA_N, 1, hi, + &valid); + } } hi = params->gamma2; - for (i = 0; valid && i < params->k; i++) { + for (i = 0; (ret == 0) && valid && (i < params->k); i++) { /* Step 25: ct0 = NTT-1(c o t0) */ - mldsa_mul_invntt(ct0 + i * MLDSA_N, c, + ret = mldsa_mul_invntt(ct0 + i * MLDSA_N, c, t0 + i * MLDSA_N); /* Step 27: Check ct0 has low enough values. */ - valid = mldsa_vec_check_low(ct0 + i * MLDSA_N, 1, hi); + if (ret == 0) { + ret = mldsa_vec_check_low(ct0 + i * MLDSA_N, 1, hi, + &valid); + } } - if (valid) { + if ((ret == 0) && valid) { /* Step 26: ct0 = ct0 + w0 */ - mldsa_vec_add(ct0, w0, params->k); - mldsa_vec_red(ct0, params->k); + ret = mldsa_vec_add(ct0, w0, params->k); + if (ret == 0) { + ret = mldsa_vec_red(ct0, params->k); + } /* Step 26, 27: Make hint from ct0 and w1 and check * number of hints is valid. * Step 32: h is encoded into signature. */ - valid = (mldsa_make_hint(ct0, w1, params->k, - params->gamma2, params->omega, h) >= 0); + if (ret == 0) { + ret = mldsa_make_hint(ct0, w1, params->k, + params->gamma2, params->omega, h, &valid); + } } } @@ -9635,7 +10076,7 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, byte* ze = sig + params->lambda / 4; /* Step 32: Encode z into signature. * Commit (c) and h already encoded into signature. */ - mldsa_vec_encode_gamma1(z, params->l, params->gamma1_bits, ze); + ret = mldsa_vec_encode_gamma1(z, params->l, params->gamma1_bits, ze); } ForceZero(priv_rand_seed, sizeof(priv_rand_seed)); @@ -9777,7 +10218,7 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, } #ifdef WOLFSSL_MLDSA_SIGN_SMALL_MEM_PRECALC if (ret == 0) { - mldsa_make_priv_vecs(key, s1, s2, t0); + ret = mldsa_make_priv_vecs(key, s1, s2, t0); } #endif #ifdef WOLFSSL_MLDSA_SIGN_SMALL_MEM_PRECALC_A @@ -9815,25 +10256,40 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, valid = 1; /* Step 12: Compute vector y from private random seed and kappa. */ - mldsa_vec_expand_mask(&key->shake, priv_rand_seed, kappa, + ret = mldsa_vec_expand_mask(&key->shake, priv_rand_seed, kappa, params->gamma1_bits, y, params->l, key->heap); #ifdef WOLFSSL_MLDSA_SIGN_CHECK_Y - valid = mldsa_vec_check_low(y, params->l, - ((sword32)1 << params->gamma1_bits) - params->beta); + if (ret == 0) { + ret = mldsa_vec_check_low(y, params->l, + ((sword32)1 << params->gamma1_bits) - params->beta, + &valid); + } #endif #ifdef WOLFSSL_MLDSA_SIGN_SMALL_MEM_PRECALC_A /* Step 13: NTT-1(A o NTT(y)) */ XMEMCPY(y_ntt, y, params->s1Sz); - mldsa_vec_ntt_full(y_ntt, params->l); - mldsa_matrix_mul(w, a, y_ntt, maxK, params->l); + if (ret == 0) { + ret = mldsa_vec_ntt_full(y_ntt, params->l); + } + if (ret == 0) { + ret = mldsa_matrix_mul(w, a, y_ntt, maxK, params->l); + } #ifdef WOLFSSL_MLDSA_SMALL - mldsa_vec_red(w, params->k); + if (ret == 0) { + ret = mldsa_vec_red(w, params->k); + } #endif - mldsa_vec_invntt_full(w, maxK); + if (ret == 0) { + ret = mldsa_vec_invntt_full(w, maxK); + } /* Step 14, Step 22: Make values positive and decompose. */ - mldsa_vec_make_pos(w, maxK); - mldsa_vec_decompose(w, maxK, params->gamma2, w0, w1); + if (ret == 0) { + ret = mldsa_vec_make_pos(w, maxK); + } + if (ret == 0) { + ret = mldsa_vec_decompose(w, maxK, params->gamma2, w0, w1); + } #endif /* Step 5: Create the matrix A from the public seed. */ /* Copy the seed into a buffer that has space for s and r. */ @@ -9882,7 +10338,12 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, break; } #endif - mldsa_ntt_full(y_ntt_t); + if (ret == 0) { + ret = mldsa_ntt_full(y_ntt_t); + } + if (ret != 0) { + break; + } /* Matrix multiply. */ #ifndef WOLFSSL_MLDSA_SMALL_MEM_POLY64 if (s == 0) { @@ -9988,9 +10449,13 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, wt[e] = mldsa_mont_red(t64[e]); } #endif - mldsa_invntt_full(wt); + if (ret == 0) { + ret = mldsa_invntt_full(wt); + } /* Step 14, Step 22: Make values positive and decompose. */ - mldsa_make_pos(wt); + if (ret == 0) { + ret = mldsa_make_pos(wt); + } #ifndef WOLFSSL_NO_ML_DSA_44 if (params->gamma2 == MLDSA_Q_LOW_88) { /* For each value of polynomial. */ @@ -10010,8 +10475,10 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, } #endif #ifdef WOLFSSL_MLDSA_SIGN_CHECK_W0 - valid = mldsa_vec_check_low(w0t, - params->gamma2 - params->beta); + if (ret == 0) { + ret = mldsa_vec_check_low(w0t, 1, + params->gamma2 - params->beta, &valid); + } #endif wt += MLDSA_N; w0t += MLDSA_N; @@ -10028,12 +10495,16 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, WC_ALLOC_VAR_EX(w1e, byte, MLDSA_MAX_W1_ENC_SZ, key->heap, DYNAMIC_TYPE_MLDSA, ret=MEMORY_E); if (WC_VAR_OK(w1e)) { - mldsa_vec_encode_w1(w1, params->k, params->gamma2, - w1e); - /* Step 15: Hash mu and encoded w1. - * Step 32: Hash is stored in signature. */ - ret = mldsa_hash256(&key->shake, mu, MLDSA_MU_SZ, - w1e, params->w1EncSz, commit, params->lambda / 4); + if (ret == 0) { + ret = mldsa_vec_encode_w1(w1, params->k, params->gamma2, + w1e); + } + if (ret == 0) { + /* Step 15: Hash mu and encoded w1. + * Step 32: Hash is stored in signature. */ + ret = mldsa_hash256(&key->shake, mu, MLDSA_MU_SZ, + w1e, params->w1EncSz, commit, params->lambda / 4); + } } WC_FREE_VAR_EX(w1e, key->heap, DYNAMIC_TYPE_MLDSA); if (ret == 0) { @@ -10044,7 +10515,7 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, } if (ret == 0) { /* Step 18: NTT(c). */ - mldsa_ntt_small(c); + ret = mldsa_ntt_small(c); } for (s = 0; (ret == 0) && valid && (s < params->l); s++) { @@ -10053,27 +10524,43 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, !defined(WOLFSSL_NO_ML_DSA_87) /* -2..2 */ if (params->eta == MLDSA_ETA_2) { - mldsa_decode_eta_2_bits(s1pt, s1); + if (ret == 0) { + ret = mldsa_decode_eta_2_bits(s1pt, s1); + } s1pt += MLDSA_ETA_2_BITS * MLDSA_N / 8; } #endif #ifndef WOLFSSL_NO_ML_DSA_65 /* -4..4 */ if (params->eta == MLDSA_ETA_4) { - mldsa_decode_eta_4_bits(s1pt, s1); + if (ret == 0) { + ret = mldsa_decode_eta_4_bits(s1pt, s1); + } s1pt += MLDSA_N / 2; } #endif - mldsa_ntt_small(s1); - mldsa_mul(z, c, s1); + if (ret == 0) { + ret = mldsa_ntt_small(s1); + } + if (ret == 0) { + ret = mldsa_mul(z, c, s1); + } #else - mldsa_mul(z, c, s1 + s * MLDSA_N); + if (ret == 0) { + ret = mldsa_mul(z, c, s1 + s * MLDSA_N); + } #endif /* Step 19: cs1 = NTT-1(c o s1) */ - mldsa_invntt(z); + if (ret == 0) { + ret = mldsa_invntt(z); + } /* Step 21: z = y + cs1 */ - mldsa_add(z, yt); - mldsa_poly_red(z); + if (ret == 0) { + ret = mldsa_add(z, yt); + } + if (ret == 0) { + ret = mldsa_poly_red(z); + } /* Step 23: Check z has low enough values. */ hi = ((sword32)1 << params->gamma1_bits) - params->beta; valid = mldsa_check_low(z, hi); @@ -10082,7 +10569,9 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, * Commit (c) and h already encoded into signature. */ #if !defined(WOLFSSL_NO_ML_DSA_44) if (params->gamma1_bits == MLDSA_GAMMA1_BITS_17) { - mldsa_encode_gamma1_17_bits(z, ze); + if (ret == 0) { + ret = mldsa_encode_gamma1_17_bits(z, ze); + } /* Move to next place to encode to. */ ze += MLDSA_GAMMA1_17_ENC_BITS / 2 * MLDSA_N / 4; @@ -10091,7 +10580,9 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, #if !defined(WOLFSSL_NO_ML_DSA_65) || \ !defined(WOLFSSL_NO_ML_DSA_87) if (params->gamma1_bits == MLDSA_GAMMA1_BITS_19) { - mldsa_encode_gamma1_19_bits(z, ze); + if (ret == 0) { + ret = mldsa_encode_gamma1_19_bits(z, ze); + } /* Move to next place to encode to. */ ze += MLDSA_GAMMA1_19_ENC_BITS / 2 * MLDSA_N / 4; @@ -10112,78 +10603,114 @@ static int mldsa_sign_with_seed_mu(wc_MlDsaKey* key, w0t = w0; w1t = w1; - for (r = 0; valid && (r < params->k); r++) { + for (r = 0; (ret == 0) && valid && (r < params->k); r++) { #ifndef WOLFSSL_MLDSA_SIGN_SMALL_MEM_PRECALC #if !defined(WOLFSSL_NO_ML_DSA_44) || \ !defined(WOLFSSL_NO_ML_DSA_87) /* -2..2 */ if (params->eta == MLDSA_ETA_2) { - mldsa_decode_eta_2_bits(s2pt, s2); + if (ret == 0) { + ret = mldsa_decode_eta_2_bits(s2pt, s2); + } s2pt += MLDSA_ETA_2_BITS * MLDSA_N / 8; } #endif #ifndef WOLFSSL_NO_ML_DSA_65 /* -4..4 */ if (params->eta == MLDSA_ETA_4) { - mldsa_decode_eta_4_bits(s2pt, s2); + if (ret == 0) { + ret = mldsa_decode_eta_4_bits(s2pt, s2); + } s2pt += MLDSA_N / 2; } #endif - mldsa_ntt_small(s2); + if (ret == 0) { + ret = mldsa_ntt_small(s2); + } /* Step 20: cs2 = NTT-1(c o s2) */ - mldsa_mul(cs2, c, s2); + if (ret == 0) { + ret = mldsa_mul(cs2, c, s2); + } #else /* Step 20: cs2 = NTT-1(c o s2) */ - mldsa_mul(cs2, c, s2 + r * MLDSA_N); + if (ret == 0) { + ret = mldsa_mul(cs2, c, s2 + r * MLDSA_N); + } #endif - mldsa_invntt(cs2); + if (ret == 0) { + ret = mldsa_invntt(cs2); + } /* Step 22: w0 - cs2 */ - mldsa_sub(w0t, cs2); - mldsa_poly_red(w0t); + if (ret == 0) { + ret = mldsa_sub(w0t, cs2); + } + if (ret == 0) { + ret = mldsa_poly_red(w0t); + } /* Step 23: Check w0 - cs2 has low enough values. */ hi = params->gamma2 - params->beta; valid = mldsa_check_low(w0t, hi); if (valid) { #ifndef WOLFSSL_MLDSA_SIGN_SMALL_MEM_PRECALC - mldsa_decode_t0(t0pt, t0); - mldsa_ntt(t0); + if (ret == 0) { + ret = mldsa_decode_t0(t0pt, t0); + } + if (ret == 0) { + ret = mldsa_ntt(t0); + } /* Step 25: ct0 = NTT-1(c o t0) */ - mldsa_mul(ct0, c, t0); + if (ret == 0) { + ret = mldsa_mul(ct0, c, t0); + } #else /* Step 25: ct0 = NTT-1(c o t0) */ - mldsa_mul(ct0, c, t0 + r * MLDSA_N); + if (ret == 0) { + ret = mldsa_mul(ct0, c, t0 + r * MLDSA_N); + } #endif - mldsa_invntt(ct0); + if (ret == 0) { + ret = mldsa_invntt(ct0); + } /* Step 27: Check ct0 has low enough values. */ valid = mldsa_check_low(ct0, params->gamma2); } if (valid) { /* Step 26: ct0 = ct0 + w0 */ - mldsa_add(ct0, w0t); - mldsa_poly_red(ct0); + if (ret == 0) { + ret = mldsa_add(ct0, w0t); + } + if (ret == 0) { + ret = mldsa_poly_red(ct0); + } /* Step 26, 27: Make hint from ct0 and w1 and check * number of hints is valid. * Step 32: h is encoded into signature. */ #ifndef WOLFSSL_NO_ML_DSA_44 - if (params->gamma2 == MLDSA_Q_LOW_88) { - valid = (mldsa_make_hint_88(ct0, w1t, h, - &idx) == 0); + if ((ret == 0) && + (params->gamma2 == MLDSA_Q_LOW_88)) { + ret = mldsa_make_hint_88(ct0, w1t, h, &idx, + &valid); /* Alg 14, Step 10: Store count of hints for * polynomial at end of list. */ - h[PARAMS_ML_DSA_44_OMEGA + r] = idx; + if (ret == 0) { + h[PARAMS_ML_DSA_44_OMEGA + r] = idx; + } } #endif #if !defined(WOLFSSL_NO_ML_DSA_65) || \ !defined(WOLFSSL_NO_ML_DSA_87) - if (params->gamma2 == MLDSA_Q_LOW_32) { - valid = (mldsa_make_hint_32(ct0, w1t, - params->omega, h, &idx) == 0); + if ((ret == 0) && + (params->gamma2 == MLDSA_Q_LOW_32)) { + ret = mldsa_make_hint_32(ct0, w1t, + params->omega, h, &idx, &valid); /* Alg 14, Step 10: Store count of hints for * polynomial at end of list. */ - h[params->omega + r] = idx; + if (ret == 0) { + h[params->omega + r] = idx; + } } #endif } @@ -10644,17 +11171,23 @@ static int mldsa_sign_ctx_hash(wc_MlDsaKey* key, WC_RNG* rng, * @param [in, out] key Key with public key data. * @param [out] t1 Vector in NTT form. */ -static void mldsa_make_pub_vec(wc_MlDsaKey* key, sword32* t1) +static int mldsa_make_pub_vec(wc_MlDsaKey* key, sword32* t1) { + int ret; const wc_MlDsaParams* params = key->params; const byte* t1p = key->p + MLDSA_PUB_SEED_SZ; - mldsa_vec_decode_t1(t1p, params->k, t1); - mldsa_vec_ntt_full(t1, params->k); + ret = mldsa_vec_decode_t1(t1p, params->k, t1); + if (ret == 0) { + ret = mldsa_vec_ntt_full(t1, params->k); + } #ifdef WC_MLDSA_CACHE_PUB_VECTORS - key->pubVecSet = 1; + if (ret == 0) { + key->pubVecSet = 1; + } #endif + return ret; } #endif @@ -10826,10 +11359,12 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, if (ret == 0) { /* Step 2: Decode z from signature. */ - mldsa_vec_decode_gamma1(ze, params->l, params->gamma1_bits, z); + ret = mldsa_vec_decode_gamma1(ze, params->l, params->gamma1_bits, z); /* Step 13: Check z is valid - values are low enough. */ hi = ((sword32)1 << params->gamma1_bits) - params->beta; - valid = mldsa_vec_check_low(z, params->l, hi); + if (ret == 0) { + ret = mldsa_vec_check_low(z, params->l, hi, &valid); + } } if ((ret == 0) && valid) { #ifdef WC_MLDSA_CACHE_PUB_VECTORS @@ -10838,7 +11373,7 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, #endif { /* Step 1: Decode and NTT vector t1. */ - mldsa_make_pub_vec(key, t1); + ret = mldsa_make_pub_vec(key, t1); } #ifdef WOLFSSL_MLDSA_VERIFY_PRECOMP_A @@ -10852,7 +11387,9 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, { #ifdef WC_MLDSA_CACHE_MATRIX_A /* Check that we haven't already cached the matrix A. */ - if (!key->aSet) + if ((ret == 0) && (!key->aSet)) +#else + if (ret == 0) #endif { /* Step 5: Expand pub seed to compute matrix A. */ @@ -10874,29 +11411,48 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, } if ((ret == 0) && valid) { /* Step 10: w = NTT-1(A o NTT(z) - NTT(c) o NTT(t1)) */ - mldsa_vec_ntt_full(z, params->l); - mldsa_matrix_mul(w, aRead, z, params->k, params->l); + ret = mldsa_vec_ntt_full(z, params->l); + if (ret == 0) { + ret = mldsa_matrix_mul(w, aRead, z, params->k, params->l); + } #ifdef WOLFSSL_MLDSA_SMALL - mldsa_vec_red(w, params->k); + if (ret == 0) { + ret = mldsa_vec_red(w, params->k); + } #endif - mldsa_ntt_small_full(c); - mldsa_vec_mul(t1c, c, t1, params->k); - mldsa_vec_sub(w, t1c, params->k); - mldsa_vec_invntt_full(w, params->k); + if (ret == 0) { + ret = mldsa_ntt_small_full(c); + } + if (ret == 0) { + ret = mldsa_vec_mul(t1c, c, t1, params->k); + } + if (ret == 0) { + ret = mldsa_vec_sub(w, t1c, params->k); + } + if (ret == 0) { + ret = mldsa_vec_invntt_full(w, params->k); + } /* Step 11: Use hint to give full w1. */ - mldsa_vec_use_hint(w, params->k, params->gamma2, params->omega, h); + if (ret == 0) { + ret = mldsa_vec_use_hint(w, params->k, params->gamma2, + params->omega, h); + } /* Step 12: Encode w1. */ - mldsa_vec_encode_w1(w, params->k, params->gamma2, w1e); - /* Step 12: Hash mu and encoded w1. */ - ret = mldsa_hash256(&key->shake, mu, MLDSA_MU_SZ, w1e, - params->w1EncSz, commit_calc, params->lambda / 4); + if (ret == 0) { + ret = mldsa_vec_encode_w1(w, params->k, params->gamma2, w1e); + } + if (ret == 0) { + /* Step 12: Hash mu and encoded w1. */ + ret = mldsa_hash256(&key->shake, mu, MLDSA_MU_SZ, w1e, + params->w1EncSz, commit_calc, params->lambda / 4); + } } if ((ret == 0) && valid) { /* Step 13: Compare commit. */ valid = (XMEMCMP(commit, commit_calc, params->lambda / 4) == 0); } - *res = valid; + *res = (ret == 0) ? valid : 0; XFREE(z, key->heap, DYNAMIC_TYPE_MLDSA); return ret; #else @@ -11001,23 +11557,30 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, unsigned int zi; valid = 1; - for (zi = 0; valid && (zi < params->l); zi++) { - mldsa_decode_gamma1(zp, params->gamma1_bits, z); + for (zi = 0; (ret == 0) && valid && (zi < params->l); zi++) { + ret = mldsa_decode_gamma1(zp, params->gamma1_bits, z); valid = mldsa_check_low(z, hi); zp += zStride; } } #else /* Step 2: Decode z from signature. */ - mldsa_vec_decode_gamma1(ze, params->l, params->gamma1_bits, z); - valid = mldsa_vec_check_low(z, params->l, hi); + if (ret == 0) { + ret = mldsa_vec_decode_gamma1(ze, params->l, + params->gamma1_bits, z); + } + if (ret == 0) { + ret = mldsa_vec_check_low(z, params->l, hi, &valid); + } #endif } if ((ret == 0) && valid) { #ifndef WOLFSSL_MLDSA_VERIFY_SMALLEST_MEM /* Step 10: NTT(z) */ - mldsa_vec_ntt_full(z, params->l); + ret = mldsa_vec_ntt_full(z, params->l); #endif + } + if ((ret == 0) && valid) { /* Step 9: Compute c from first 256 bits of commit. */ #ifdef WOLFSSL_MLDSA_VERIFY_NO_MALLOC @@ -11029,7 +11592,7 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, #endif } if ((ret == 0) && valid) { - mldsa_ntt_small_full(c); + ret = mldsa_ntt_small_full(c); o = 0; encW1 = w1e; @@ -11043,12 +11606,16 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, const sword32* zt = z; /* Step 1: Decode and NTT vector t1. */ - mldsa_decode_t1(t1p, w); + if (ret == 0) { + ret = mldsa_decode_t1(t1p, w); + } /* Next polynomial. */ t1p += MLDSA_U * MLDSA_N / 8; /* Step 10: - NTT(c) o NTT(t1)) */ - mldsa_ntt_full(w); + if (ret == 0) { + ret = mldsa_ntt_full(w); + } #ifndef WOLFSSL_MLDSA_SMALL_MEM_POLY64 #ifdef WOLFSSL_MLDSA_SMALL for (e = 0; e < MLDSA_N; e++) { @@ -11094,9 +11661,13 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, #ifdef WOLFSSL_MLDSA_VERIFY_SMALLEST_MEM /* Step 2/10: Decode and NTT this z polynomial on demand (z is * not kept as a whole vector in this mode). */ - mldsa_decode_gamma1(ze + (word32)s * zStride, - params->gamma1_bits, z); - mldsa_ntt_full(z); + if (ret == 0) { + ret = mldsa_decode_gamma1(ze + (word32)s * zStride, + params->gamma1_bits, z); + } + if (ret == 0) { + ret = mldsa_ntt_full(z); + } zt = z; #endif /* Step 3: Create polynomial from hashing seed. */ @@ -11111,11 +11682,15 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, #endif { #ifdef WOLFSSL_MLDSA_VERIFY_NO_MALLOC - ret = mldsa_rej_ntt_poly_ex(&key->shake, seed, aBuf, - key->h); + if (ret == 0) { + ret = mldsa_rej_ntt_poly_ex(&key->shake, seed, aBuf, + key->h); + } #else - ret = mldsa_rej_ntt_poly_ex(&key->shake, seed, aBuf, - block); + if (ret == 0) { + ret = mldsa_rej_ntt_poly_ex(&key->shake, seed, aBuf, + block); + } #endif a = aBuf; } @@ -11166,14 +11741,18 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, #endif /* Step 10: w = NTT-1(A o NTT(z) - NTT(c) o NTT(t1)) */ - mldsa_invntt_full(w); + if (ret == 0) { + ret = mldsa_invntt_full(w); + } #ifndef WOLFSSL_NO_ML_DSA_44 if (params->gamma2 == MLDSA_Q_LOW_88) { /* Step 11: Use hint to give full w1. */ mldsa_use_hint_88(w, h, r, &o); /* Step 12: Encode w1. */ - mldsa_encode_w1_88(w, encW1); + if (ret == 0) { + ret = mldsa_encode_w1_88(w, encW1); + } encW1 += MLDSA_Q_HI_88_ENC_BITS * 2 * MLDSA_N / 16; } else @@ -11183,7 +11762,9 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, /* Step 11: Use hint to give full w1. */ mldsa_use_hint_32(w, h, params->omega, r, &o); /* Step 12: Encode w1. */ - mldsa_encode_w1_32(w, encW1); + if (ret == 0) { + ret = mldsa_encode_w1_32(w, encW1); + } encW1 += MLDSA_Q_HI_32_ENC_BITS * 2 * MLDSA_N / 16; } else @@ -11202,7 +11783,7 @@ static int mldsa_verify_with_mu(wc_MlDsaKey* key, const byte* mu, valid = (XMEMCMP(commit, commit_calc, params->lambda / 4) == 0); } - *res = valid; + *res = (ret == 0) ? valid : 0; #ifndef WOLFSSL_MLDSA_VERIFY_NO_MALLOC XFREE(z, key->heap, DYNAMIC_TYPE_MLDSA); #endif @@ -12765,8 +13346,12 @@ int wc_MlDsaKey_CheckKey(wc_MlDsaKey* key) sword32 x = 0; /* Get s1, s2 and t0 from private key. */ - mldsa_vec_decode_eta_bits(s1p, params->eta, s1, params->l); - mldsa_vec_decode_eta_bits(s2p, params->eta, s2, params->k); + if (ret == 0) { + ret = mldsa_vec_decode_eta_bits(s1p, params->eta, s1, params->l); + } + if (ret == 0) { + ret = mldsa_vec_decode_eta_bits(s2p, params->eta, s2, params->k); + } /* Validate s1 and s2 coefficients are within [-eta, eta]. */ { sword32 eta = (sword32)params->eta; @@ -12785,23 +13370,39 @@ int wc_MlDsaKey_CheckKey(wc_MlDsaKey* key) } } if (ret == 0) { - mldsa_vec_decode_t0(t0p, params->k, t0); + ret = mldsa_vec_decode_t0(t0p, params->k, t0); /* Get t1 from public key. */ - mldsa_vec_decode_t1(t1p, params->k, t1); + if (ret == 0) { + ret = mldsa_vec_decode_t1(t1p, params->k, t1); + } /* Calculate t = NTT-1(A o NTT(s1)) + s2 */ - mldsa_vec_ntt_small_full(s1, params->l); - mldsa_matrix_mul(t, a, s1, params->k, params->l); + if (ret == 0) { + ret = mldsa_vec_ntt_small_full(s1, params->l); + } + if (ret == 0) { + ret = mldsa_matrix_mul(t, a, s1, params->k, params->l); + } #ifdef WOLFSSL_MLDSA_SMALL - mldsa_vec_red(t, params->k); + if (ret == 0) { + ret = mldsa_vec_red(t, params->k); + } #endif - mldsa_vec_invntt_full(t, params->k); - mldsa_vec_add(t, s2, params->k); + if (ret == 0) { + ret = mldsa_vec_invntt_full(t, params->k); + } + if (ret == 0) { + ret = mldsa_vec_add(t, s2, params->k); + } /* Subtract t0 from t. */ - mldsa_vec_sub(t, t0, params->k); + if (ret == 0) { + ret = mldsa_vec_sub(t, t0, params->k); + } /* Make t positive to match t1. */ - mldsa_vec_make_pos(t, params->k); + if (ret == 0) { + ret = mldsa_vec_make_pos(t, params->k); + } /* Check t - t0 and t1 are the same. */ for (i = 0; i < params->k; i++) { @@ -12816,7 +13417,7 @@ int wc_MlDsaKey_CheckKey(wc_MlDsaKey* key) x |= key->p[i] ^ key->k[i]; } - if (x != 0) { + if ((ret == 0) && (x != 0)) { ret = PUBLIC_KEY_E; } } @@ -13006,6 +13607,14 @@ int wc_MlDsaKey_ImportPubRaw(wc_MlDsaKey* key, const byte* in, word32 inLen) #endif if (ret == 0) { + /* A failed rebuild below must not leave the old key usable. */ + key->pubKeySet = 0; + #ifdef WC_MLDSA_CACHE_PUB_VECTORS + key->pubVecSet = 0; + #endif + #ifdef WC_MLDSA_CACHE_MATRIX_A + key->aSet = 0; + #endif /* Copy the private key data in or copy pointer. */ #ifdef WOLFSSL_MLDSA_ASSIGN_KEY key->p = in; @@ -13030,7 +13639,7 @@ int wc_MlDsaKey_ImportPubRaw(wc_MlDsaKey* key, const byte* in, word32 inLen) } if (ret == 0) { /* Compute t1 from public key data. */ - mldsa_make_pub_vec(key, key->t1); + ret = mldsa_make_pub_vec(key, key->t1); #endif #ifdef WC_MLDSA_CACHE_MATRIX_A #ifndef WC_MLDSA_FIXED_ARRAY @@ -13165,6 +13774,14 @@ static int mldsa_set_priv_key(const byte* priv, word32 privSz, } if (ret == 0) { + /* A failed rebuild below must not leave the old key usable. */ + key->prvKeySet = 0; + #ifdef WC_MLDSA_CACHE_PRIV_VECTORS + key->privVecsSet = 0; + #endif + #ifdef WC_MLDSA_CACHE_MATRIX_A + key->aSet = 0; + #endif /* Copy the private key data in or copy pointer. */ #ifdef WOLFSSL_MLDSA_ASSIGN_KEY key->k = priv; @@ -13221,7 +13838,7 @@ static int mldsa_set_priv_key(const byte* priv, word32 privSz, #endif if (ret == 0) { /* Compute vectors from private key. */ - mldsa_make_priv_vecs(key, key->s1, key->s2, key->t0); + ret = mldsa_make_priv_vecs(key, key->s1, key->s2, key->t0); } #endif if (ret == 0) { diff --git a/wolfcrypt/src/wc_mlkem.c b/wolfcrypt/src/wc_mlkem.c index e9e582a88b5..6328ab531e1 100644 --- a/wolfcrypt/src/wc_mlkem.c +++ b/wolfcrypt/src/wc_mlkem.c @@ -1348,15 +1348,17 @@ static int mlkemkey_encapsulate(MlKemKey* key, const byte* m, byte* r, byte* c) /* Convert msg to a polynomial. * Step 20: mu <- Decompress_1(ByteDecode_1(m)) */ - MLKEM_ARM64_SVR(mlkem_from_msg(mu, m)); + MLKEM_ARM64_SVR(ret = mlkem_from_msg(mu, m)); } if (ret == 0) { /* Initialize the PRF for use in the noise generation. */ mlkem_prf_init(&key->prf); - /* Generate noise using PRF. - * Steps 9-17: generate y, e_1, e_2 - */ - ret = mlkem_get_noise(&key->prf, (int)k, y, e1, e2, r); + if (ret == 0) { + /* Generate noise using PRF. + * Steps 9-17: generate y, e_1, e_2 + */ + ret = mlkem_get_noise(&key->prf, (int)k, y, e1, e2, r); + } } #ifdef WOLFSSL_MLKEM_CACHE_A if ((ret == 0) && ((key->flags & MLKEM_FLAG_A_SET) != 0)) { @@ -1424,8 +1426,9 @@ static int mlkemkey_encapsulate(MlKemKey* key, const byte* m, byte* r, byte* c) /* Step 22: c_1 <- ByteEncode_d_u(Compress_d_u(u)) * Step 23: c_2 <- ByteEncode_d_v(Compress_d_v(v)) */ MLKEM_ARM64_SVR({ - mlkem_vec_compress_10(c1, u, k); - mlkem_compress_4(c2, v); + ret = mlkem_vec_compress_10(c1, u, k); + if (ret == 0) + ret = mlkem_compress_4(c2, v); }); /* Step 24: return c <- (c_1||c_2) */ } @@ -1435,8 +1438,9 @@ static int mlkemkey_encapsulate(MlKemKey* key, const byte* m, byte* r, byte* c) /* Step 22: c_1 <- ByteEncode_d_u(Compress_d_u(u)) * Step 23: c_2 <- ByteEncode_d_v(Compress_d_v(v)) */ MLKEM_ARM64_SVR({ - mlkem_vec_compress_10(c1, u, k); - mlkem_compress_4(c2, v); + ret = mlkem_vec_compress_10(c1, u, k); + if (ret == 0) + ret = mlkem_compress_4(c2, v); }); /* Step 24: return c <- (c_1||c_2) */ } @@ -1446,8 +1450,9 @@ static int mlkemkey_encapsulate(MlKemKey* key, const byte* m, byte* r, byte* c) /* Step 22: c_1 <- ByteEncode_d_u(Compress_d_u(u)) * Step 23: c_2 <- ByteEncode_d_v(Compress_d_v(v)) */ MLKEM_ARM64_SVR({ - mlkem_vec_compress_11(c1, u); - mlkem_compress_5(c2, v); + ret = mlkem_vec_compress_11(c1, u); + if (ret == 0) + ret = mlkem_compress_5(c2, v); }); /* Step 24: return c <- (c_1||c_2) */ } @@ -1937,37 +1942,51 @@ static MLKEM_NOINLINE int mlkemkey_decapsulate(MlKemKey* key, byte* m, #if defined(WOLFSSL_KYBER512) || defined(WOLFSSL_WC_ML_KEM_512) if (k == WC_ML_KEM_512_K) { /* Step 3: u' <- Decompress_d_u(ByteDecode_d_u(c1)) */ - mlkem_vec_decompress_10(u, c1, k); + if (ret == 0) { + ret = mlkem_vec_decompress_10(u, c1, k); + } /* Step 4: v' <- Decompress_d_v(ByteDecode_d_v(c2)) */ - mlkem_decompress_4(v, c2); + if (ret == 0) { + ret = mlkem_decompress_4(v, c2); + } } #endif #if defined(WOLFSSL_KYBER768) || defined(WOLFSSL_WC_ML_KEM_768) if (k == WC_ML_KEM_768_K) { /* Step 3: u' <- Decompress_d_u(ByteDecode_d_u(c1)) */ - mlkem_vec_decompress_10(u, c1, k); + if (ret == 0) { + ret = mlkem_vec_decompress_10(u, c1, k); + } /* Step 4: v' <- Decompress_d_v(ByteDecode_d_v(c2)) */ - mlkem_decompress_4(v, c2); + if (ret == 0) { + ret = mlkem_decompress_4(v, c2); + } } #endif #if defined(WOLFSSL_KYBER1024) || defined(WOLFSSL_WC_ML_KEM_1024) if (k == WC_ML_KEM_1024_K) { /* Step 3: u' <- Decompress_d_u(ByteDecode_d_u(c1)) */ - mlkem_vec_decompress_11(u, c1); + if (ret == 0) { + ret = mlkem_vec_decompress_11(u, c1); + } /* Step 4: v' <- Decompress_d_v(ByteDecode_d_v(c2)) */ - mlkem_decompress_5(v, c2); + if (ret == 0) { + ret = mlkem_decompress_5(v, c2); + } } #endif /* Decapsulate the cipher text into polynomial. * Step 6: w <- v' - InvNTT(s_hat_trans o NTT(u')) */ - ret = mlkem_decapsulate(key->priv, w, u, v, (int)k); + if (ret == 0) { + ret = mlkem_decapsulate(key->priv, w, u, v, (int)k); + } } if (ret == 0) { /* Convert the polynomial into a array of bytes (message). * Step 7: m <- ByteEncode_1(Compress_1(w)) */ - MLKEM_ARM64_SVR(mlkem_to_msg(m, w)); + MLKEM_ARM64_SVR(ret = mlkem_to_msg(m, w)); /* Step 8: return m */ } @@ -2127,6 +2146,11 @@ int wc_MlKemKey_Decapsulate(MlKemKey* key, unsigned char* ss, ret = 0; } #endif + /* Software decapsulation re-encrypts with the public key (FIPS 203, + * Algorithm 18, steps 2 and 8), so refuse before decrypting without it. */ + if ((ret == 0) && ((key->flags & MLKEM_FLAG_PUB_SET) == 0)) { + ret = BAD_STATE_E; + } #if !defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_NO_MALLOC) if (ret == 0) { @@ -2166,7 +2190,7 @@ int wc_MlKemKey_Decapsulate(MlKemKey* key, unsigned char* ss, } if (ret == 0) { /* Compare generated cipher text with that passed in. */ - MLKEM_ARM64_SVR(fail = mlkem_cmp(ct, cmp, (int)ctSz)); + MLKEM_ARM64_SVR(ret = mlkem_cmp(ct, cmp, (int)ctSz, &fail)); } if (ret == 0) { #if defined(WOLFSSL_MLKEM_KYBER) && !defined(WOLFSSL_NO_ML_KEM) @@ -2243,15 +2267,17 @@ int wc_MlKemKey_Decapsulate(MlKemKey* key, unsigned char* ss, * @param [out] pubSeed Public seed. * @param [in] p Public key data. * @param [in] k Number of polynomials in vector. + * @return 0 on success, or the error from a refused vector-register save. */ -static void mlkemkey_decode_public(sword16* pub, byte* pubSeed, const byte* p, +static int mlkemkey_decode_public(sword16* pub, byte* pubSeed, const byte* p, unsigned int k) { + int ret; unsigned int i; /* Decode public key that is vector of polynomials. * Step 2: t <- ByteDecode_12(ek_PKE[0 : 384k]) */ - mlkem_from_bytes(pub, p, (int)k); + ret = mlkem_from_bytes(pub, p, (int)k); p += k * WC_ML_KEM_POLY_SIZE; /* Read public key seed. @@ -2259,6 +2285,7 @@ static void mlkemkey_decode_public(sword16* pub, byte* pubSeed, const byte* p, for (i = 0; i < WC_ML_KEM_SYM_SZ; i++) { pubSeed[i] = p[i]; } + return ret; } /** @@ -2367,6 +2394,12 @@ int wc_MlKemKey_DecodePrivateKey(MlKemKey* key, const unsigned char* in, if ((ret == 0) && (len != privLen)) { ret = BUFFER_E; } + if (ret == 0) { + /* Forget the old key before its buffers are replaced, so a failure + * from here on leaves the key unusable. */ + key->flags &= ~(MLKEM_FLAG_BOTH_SET | MLKEM_FLAG_H_SET | + MLKEM_FLAG_A_SET); + } #ifdef WOLFSSL_MLKEM_DYNAMIC_KEYS if (ret == 0) { @@ -2377,22 +2410,21 @@ int wc_MlKemKey_DecodePrivateKey(MlKemKey* key, const unsigned char* in, } #endif if (ret == 0) { - /* Clear the key-set flags first so any failure below (size, reduction - * check, or hash) leaves a reused key object consistently unusable - * rather than flagged-set with zeroed material. */ - key->flags &= ~(MLKEM_FLAG_BOTH_SET | MLKEM_FLAG_H_SET); - /* Decode private key that is vector of polynomials. * Alg 18 Step 1: dk_PKE <- dk[0 : 384k] * Alg 15 Step 5: s_hat <- ByteDecode_12(dk_PKE) */ - mlkem_from_bytes(key->priv, p, (int)k); + ret = mlkem_from_bytes(key->priv, p, (int)k); p += k * WC_ML_KEM_POLY_SIZE; /* Both vectors must decode to coefficients reduced modulo q. */ - ret = mlkem_check_reduced(key->priv, (int)k); + if (ret == 0) { + ret = mlkem_check_reduced(key->priv, (int)k); + } if (ret == 0) { /* Decode the public key that is after the private key. */ - mlkemkey_decode_public(key->pub, key->pubSeed, p, k); + ret = mlkemkey_decode_public(key->pub, key->pubSeed, p, k); + } + if (ret == 0) { ret = mlkem_check_reduced(key->pub, (int)k); } if (ret != 0) { @@ -2510,6 +2542,12 @@ int wc_MlKemKey_DecodePublicKey(MlKemKey* key, const unsigned char* in, if ((ret == 0) && (len != pubLen)) { ret = BUFFER_E; } + if (ret == 0) { + /* Forget the old public key and its cached matrix before they are + * replaced, so a failure from here on leaves no public key. */ + key->flags &= ~(MLKEM_FLAG_PUB_SET | MLKEM_FLAG_H_SET | + MLKEM_FLAG_A_SET); + } #ifdef WOLFSSL_MLKEM_DYNAMIC_KEYS if (ret == 0) { @@ -2518,7 +2556,9 @@ int wc_MlKemKey_DecodePublicKey(MlKemKey* key, const unsigned char* in, #endif if (ret == 0) { /* Decode public key and check public key matches parameters. */ - mlkemkey_decode_public(key->pub, key->pubSeed, p, k); + ret = mlkemkey_decode_public(key->pub, key->pubSeed, p, k); + } + if (ret == 0) { ret = mlkem_check_reduced(key->pub, (int)k); } if (ret == 0) { @@ -2764,7 +2804,7 @@ int wc_MlKemKey_EncodePrivateKey(MlKemKey* key, unsigned char* out, word32 len) if (ret == 0) { /* Encode private key that is vector of polynomials. */ - MLKEM_ARM64_SVR(mlkem_to_bytes(p, key->priv, (int)k)); + MLKEM_ARM64_SVR(ret = mlkem_to_bytes(p, key->priv, (int)k)); } if (ret == 0) { p += WC_ML_KEM_POLY_SIZE * k; @@ -2876,7 +2916,7 @@ int wc_MlKemKey_EncodePublicKey(MlKemKey* key, unsigned char* out, word32 len) if (ret == 0) { /* Encode public key polynomial by polynomial. */ - MLKEM_ARM64_SVR(mlkem_to_bytes(p, key->pub, (int)k)); + MLKEM_ARM64_SVR(ret = mlkem_to_bytes(p, key->pub, (int)k)); } if (ret == 0) { int i; diff --git a/wolfcrypt/src/wc_mlkem_poly.c b/wolfcrypt/src/wc_mlkem_poly.c index 616da2c3843..7f4e5b4ad17 100644 --- a/wolfcrypt/src/wc_mlkem_poly.c +++ b/wolfcrypt/src/wc_mlkem_poly.c @@ -1948,14 +1948,20 @@ int mlkem_keygen(sword16* s, sword16* t, sword16* e, const sword16* a, int k) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; /* Alg 13: Steps 16-18 */ mlkem_keygen_avx512(s, t, e, a, k); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; /* Alg 13: Steps 16-18 */ mlkem_keygen_avx2(s, t, e, a, k); RESTORE_VECTOR_REGISTERS(); @@ -2163,13 +2169,19 @@ int mlkem_encapsulate(const sword16* pub, sword16* u, sword16* v, { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_encapsulate_avx512(pub, u, v, a, y, e1, e2, m, k); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_encapsulate_avx2(pub, u, v, a, y, e1, e2, m, k); RESTORE_VECTOR_REGISTERS(); } @@ -2269,11 +2281,15 @@ int mlkem_encapsulate_seeds(const sword16* pub, MLKEM_PRF_T* prf, sword16* u, /* Inverse transform v. */ mlkem_invntt(v); - mlkem_from_msg(m, msg); + if (ret == 0) { + ret = mlkem_from_msg(m, msg); + } /* Generate noise using PRF. */ coins[WC_ML_KEM_SYM_SZ] = WC_OCTET(2 * k); - ret = mlkem_get_noise_eta2_c(prf, e2, coins); + if (ret == 0) { + ret = mlkem_get_noise_eta2_c(prf, e2, coins); + } if (ret == 0) { /* Add errors and message to v and reduce. */ #if defined(WOLFSSL_MLKEM_SMALL) || defined(WOLFSSL_MLKEM_NO_LARGE_CODE) @@ -2368,13 +2384,19 @@ int mlkem_decapsulate(const sword16* s, sword16* w, sword16* u, { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decapsulate_avx512(s, w, u, v, k); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decapsulate_avx2(s, w, u, v, k); RESTORE_VECTOR_REGISTERS(); } @@ -2821,8 +2843,13 @@ static int mlkem_gen_matrix_k3_avx2(sword16* a, byte* seed, int transposed) if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + WC_FREE_VAR_EX(rand, NULL, DYNAMIC_TYPE_TMP_BUFFER); + WC_FREE_VAR_EX(state, NULL, DYNAMIC_TYPE_TMP_BUFFER); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -2839,8 +2866,13 @@ static int mlkem_gen_matrix_k3_avx2(sword16* a, byte* seed, int transposed) if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + WC_FREE_VAR_EX(rand, NULL, DYNAMIC_TYPE_TMP_BUFFER); + WC_FREE_VAR_EX(state, NULL, DYNAMIC_TYPE_TMP_BUFFER); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -2949,8 +2981,13 @@ static int mlkem_gen_matrix_k3_avx512(sword16* a, byte* seed, int transposed) if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + WC_FREE_VAR_EX(rand, NULL, DYNAMIC_TYPE_TMP_BUFFER); + WC_FREE_VAR_EX(state, NULL, DYNAMIC_TYPE_TMP_BUFFER); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -2967,8 +3004,13 @@ static int mlkem_gen_matrix_k3_avx512(sword16* a, byte* seed, int transposed) if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) - { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + WC_FREE_VAR_EX(rand, NULL, DYNAMIC_TYPE_TMP_BUFFER); + WC_FREE_VAR_EX(state, NULL, DYNAMIC_TYPE_TMP_BUFFER); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -3420,7 +3462,7 @@ static int mlkem_gen_matrix_k4_aarch64(sword16* a, byte* seed, int transposed) * @param [in, out] shake128 SHAKE-128 object. * @param [in] seed Data to absorb. * @param [in] len Length of data to absorb in bytes. - * @return 0 on success always. + * @return 0 on success, or the error from a refused vector-register save. */ static int mlkem_xof_absorb(wc_Shake* shake128, const byte* seed, int len) { @@ -3582,7 +3624,7 @@ void mlkem_prf_free(wc_Shake* prf) * @param [in] outLen Number of bytes to write. * @param [in] key Data to derive from. Must be: * WC_ML_KEM_SYM_SZ + 1 bytes in length. - * @return 0 on success always. + * @return 0 on success, or the error from a refused vector-register save. */ static int mlkem_prf(wc_Shake* shake256, byte* out, unsigned int outLen, const byte* key) @@ -3612,8 +3654,12 @@ static int mlkem_prf(wc_Shake* shake256, byte* out, unsigned int outLen, if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + ForceZero(state, sizeof(state)); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -3662,7 +3708,7 @@ static int mlkem_prf(wc_Shake* shake256, byte* out, unsigned int outLen, * @param [in] seedLen Length of data to derive from in bytes. * @param [out] out Buffer to write to. * @param [in] outLen Number of bytes to derive. - * @return 0 on success always. + * @return 0 on success, or the error from a refused vector-register save. */ int mlkem_kdf(const byte* seed, int seedLen, byte* out, int outLen) { @@ -3678,7 +3724,12 @@ int mlkem_kdf(const byte* seed, int seedLen, byte* out, int outLen) if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + ForceZero(state, sizeof(state)); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -4108,13 +4159,19 @@ int mlkem_gen_matrix(MLKEM_PRF_T* prf, sword16* a, int k, byte* seed, #else #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_gen_matrix_k2_avx512(a, seed, transposed); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_gen_matrix_k2_avx2(a, seed, transposed); RESTORE_VECTOR_REGISTERS(); } @@ -4139,13 +4196,19 @@ int mlkem_gen_matrix(MLKEM_PRF_T* prf, sword16* a, int k, byte* seed, #else #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_gen_matrix_k3_avx512(a, seed, transposed); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_gen_matrix_k3_avx2(a, seed, transposed); RESTORE_VECTOR_REGISTERS(); } @@ -4170,13 +4233,19 @@ int mlkem_gen_matrix(MLKEM_PRF_T* prf, sword16* a, int k, byte* seed, #else #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_gen_matrix_k4_avx512(a, seed, transposed); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_gen_matrix_k4_avx2(a, seed, transposed); RESTORE_VECTOR_REGISTERS(); } @@ -4820,7 +4889,12 @@ static int mlkem_get_noise_eta2_avx2(MLKEM_PRF_T* prf, sword16* p, if (IS_INTEL_BMI2(cpuid_flags)) { sha3_block_bmi2(state); } - else if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + else if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + ForceZero(state, sizeof(state)); + return svr_ret; + } sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); } @@ -5540,13 +5614,19 @@ int mlkem_get_noise(MLKEM_PRF_T* prf, int k, sword16* vec1, sword16* vec2, #else #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_get_noise_k2_avx512(prf, vec1, vec2, poly, seed); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_get_noise_k2_avx2(prf, vec1, vec2, poly, seed); RESTORE_VECTOR_REGISTERS(); } @@ -5576,13 +5656,19 @@ int mlkem_get_noise(MLKEM_PRF_T* prf, int k, sword16* vec1, sword16* vec2, #else #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_get_noise_k3_avx512(vec1, vec2, poly, seed); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_get_noise_k3_avx2(vec1, vec2, poly, seed); RESTORE_VECTOR_REGISTERS(); } @@ -5608,13 +5694,19 @@ int mlkem_get_noise(MLKEM_PRF_T* prf, int k, sword16* vec1, sword16* vec2, #else #if defined(USE_INTEL_SPEEDUP) && !defined(WC_SHA3_NO_ASM) #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_get_noise_k4_avx512(prf, vec1, vec2, poly, seed); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = mlkem_get_noise_k4_avx2(prf, vec1, vec2, poly, seed); RESTORE_VECTOR_REGISTERS(); } @@ -5702,38 +5794,47 @@ static int mlkem_cmp_c(const byte* a, const byte* b, int sz) /* Compare two byte arrays of equal size. * - * @param [in] a First array to compare. - * @param [in] b Second array to compare. - * @param [in] sz Size of arrays in bytes. + * @param [in] a First array to compare. + * @param [in] b Second array to compare. + * @param [in] sz Size of arrays in bytes. + * @param [out] fail 0 when the arrays match, -1 when they differ. * @return 0 on success. - * @return -1 on failure. + * @return Error from a refused vector-register save. */ -int mlkem_cmp(const byte* a, const byte* b, int sz) +int mlkem_cmp(const byte* a, const byte* b, int sz, int* fail) { + /* Start at "did not match" so an error return cannot read as a match. */ + *fail = -1; + #if defined(__aarch64__) && defined(WOLFSSL_ARMASM) - return mlkem_cmp_neon(a, b, sz); + *fail = mlkem_cmp_neon(a, b, sz); + return 0; #else - int fail; - #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - fail = mlkem_cmp_avx512(a, b, sz); + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; + *fail = mlkem_cmp_avx512(a, b, sz); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { - fail = mlkem_cmp_avx2(a, b, sz); + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; + *fail = mlkem_cmp_avx2(a, b, sz); RESTORE_VECTOR_REGISTERS(); } else #endif { - fail = mlkem_cmp_c(a, b, sz); + *fail = mlkem_cmp_c(a, b, sz); } - return fail; + return 0; #endif } @@ -6011,26 +6112,34 @@ static void mlkem_vec_compress_10_c(byte* r, sword16* v, unsigned int k) * @param [in, out] v Vector of polynomials. * @param [in] k Number of polynomials in vector. */ -void mlkem_vec_compress_10(byte* r, sword16* v, unsigned int k) +int mlkem_vec_compress_10(byte* r, sword16* v, unsigned int k) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_10_avx512_vbmi(r, v, (int)k); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_10_avx512(r, v, (int)k); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_10_avx2(r, v, (int)k); RESTORE_VECTOR_REGISTERS(); } @@ -6039,6 +6148,7 @@ void mlkem_vec_compress_10(byte* r, sword16* v, unsigned int k) { mlkem_vec_compress_10_c(r, v, k); } + return 0; } #endif @@ -6125,17 +6235,23 @@ static void mlkem_vec_compress_11_c(byte* r, sword16* v) * @param [out] r Array of bytes. * @param [in, out] v Vector of polynomials. */ -void mlkem_vec_compress_11(byte* r, sword16* v) +int mlkem_vec_compress_11(byte* r, sword16* v) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_11_avx512(r, v, 4); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_11_avx2(r, v, 4); RESTORE_VECTOR_REGISTERS(); } @@ -6144,6 +6260,7 @@ void mlkem_vec_compress_11(byte* r, sword16* v) { mlkem_vec_compress_11_c(r, v); } + return 0; } #endif #endif /* !WOLFSSL_MLKEM_NO_ENCAPSULATE || !WOLFSSL_MLKEM_NO_DECAPSULATE */ @@ -6238,26 +6355,34 @@ static void mlkem_vec_decompress_10_c(sword16* v, const byte* b, unsigned int k) * @param [in] b Array of bytes. * @param [in] k Number of polynomials in vector. */ -void mlkem_vec_decompress_10(sword16* v, const byte* b, unsigned int k) +int mlkem_vec_decompress_10(sword16* v, const byte* b, unsigned int k) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_10_avx512_vbmi(v, b, (int)k); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_10_avx512(v, b, (int)k); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_10_avx2(v, b, (int)k); RESTORE_VECTOR_REGISTERS(); } @@ -6266,6 +6391,7 @@ void mlkem_vec_decompress_10(sword16* v, const byte* b, unsigned int k) { mlkem_vec_decompress_10_c(v, b, k); } + return 0; } #endif #if defined(WOLFSSL_KYBER1024) || defined(WOLFSSL_WC_ML_KEM_1024) @@ -6342,26 +6468,34 @@ static void mlkem_vec_decompress_11_c(sword16* v, const byte* b) * @param [out] v Vector of polynomials. * @param [in] b Array of bytes. */ -void mlkem_vec_decompress_11(sword16* v, const byte* b) +int mlkem_vec_decompress_11(sword16* v, const byte* b) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_11_avx512_vbmi(v, b, 4); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_11_avx512(v, b, 4); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_11_avx2(v, b, 4); RESTORE_VECTOR_REGISTERS(); } @@ -6370,6 +6504,7 @@ void mlkem_vec_decompress_11(sword16* v, const byte* b) { mlkem_vec_decompress_11_c(v, b); } + return 0; } #endif #endif /* !WOLFSSL_MLKEM_NO_DECAPSULATE */ @@ -6527,26 +6662,34 @@ static void mlkem_compress_4_c(byte* b, sword16* p) * @param [out] b Array of bytes. * @param [in, out] p Polynomial. */ -void mlkem_compress_4(byte* b, sword16* p) +int mlkem_compress_4(byte* b, sword16* p) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_4_avx512_vbmi(b, p); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_4_avx512(b, p); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_4_avx2(b, p); RESTORE_VECTOR_REGISTERS(); } @@ -6555,6 +6698,7 @@ void mlkem_compress_4(byte* b, sword16* p) { mlkem_compress_4_c(b, p); } + return 0; } #endif #if defined(WOLFSSL_KYBER1024) || defined(WOLFSSL_WC_ML_KEM_1024) @@ -6621,26 +6765,34 @@ static void mlkem_compress_5_c(byte* b, sword16* p) * @param [out] b Array of bytes. * @param [in, out] p Polynomial. */ -void mlkem_compress_5(byte* b, sword16* p) +int mlkem_compress_5(byte* b, sword16* p) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_5_avx512_vbmi(b, p); RESTORE_VECTOR_REGISTERS(); } else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_5_avx512(b, p); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_compress_5_avx2(b, p); RESTORE_VECTOR_REGISTERS(); } @@ -6649,6 +6801,7 @@ void mlkem_compress_5(byte* b, sword16* p) { mlkem_compress_5_c(b, p); } + return 0; } #endif #endif /* !WOLFSSL_MLKEM_NO_ENCAPSULATE || !WOLFSSL_MLKEM_NO_DECAPSULATE */ @@ -6707,17 +6860,23 @@ static void mlkem_decompress_4_c(sword16* p, const byte* b) * @param [out] p Polynomial. * @param [in] b Array of bytes. */ -void mlkem_decompress_4(sword16* p, const byte* b) +int mlkem_decompress_4(sword16* p, const byte* b) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_4_avx512(p, b); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_4_avx2(p, b); RESTORE_VECTOR_REGISTERS(); } @@ -6726,6 +6885,7 @@ void mlkem_decompress_4(sword16* p, const byte* b) { mlkem_decompress_4_c(p, b); } + return 0; } #endif #if defined(WOLFSSL_KYBER1024) || defined(WOLFSSL_WC_ML_KEM_1024) @@ -6793,17 +6953,23 @@ static void mlkem_decompress_5_c(sword16* p, const byte* b) * @param [out] p Polynomial. * @param [in] b Array of bytes. */ -void mlkem_decompress_5(sword16* p, const byte* b) +int mlkem_decompress_5(sword16* p, const byte* b) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_5_avx512(p, b); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_decompress_5_avx2(p, b); RESTORE_VECTOR_REGISTERS(); } @@ -6812,6 +6978,7 @@ void mlkem_decompress_5(sword16* p, const byte* b) { mlkem_decompress_5_c(p, b); } + return 0; } #endif #endif /* !WOLFSSL_MLKEM_NO_DECAPSULATE */ @@ -6877,17 +7044,23 @@ static void mlkem_from_msg_c(sword16* p, const byte* msg) * @param [out] p Polynomial. * @param [in] msg Message as a byte array. */ -void mlkem_from_msg(sword16* p, const byte* msg) +int mlkem_from_msg(sword16* p, const byte* msg) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_from_msg_avx512(p, msg); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; mlkem_from_msg_avx2(p, msg); RESTORE_VECTOR_REGISTERS(); } @@ -6896,6 +7069,7 @@ void mlkem_from_msg(sword16* p, const byte* msg) { mlkem_from_msg_c(p, msg); } + return 0; } #endif @@ -6985,18 +7159,24 @@ static void mlkem_to_msg_c(byte* msg, sword16* p) * @param [out] msg Message as a byte array. * @param [in, out] p Polynomial. */ -void mlkem_to_msg(byte* msg, sword16* p) +int mlkem_to_msg(byte* msg, sword16* p) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; /* Convert the polynomial into an array of bytes (message). */ mlkem_to_msg_avx512(msg, p); RESTORE_VECTOR_REGISTERS(); } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; /* Convert the polynomial into an array of bytes (message). */ mlkem_to_msg_avx2(msg, p); RESTORE_VECTOR_REGISTERS(); @@ -7006,6 +7186,7 @@ void mlkem_to_msg(byte* msg, sword16* p) { mlkem_to_msg_c(msg, p); } + return 0; } #endif /* !WOLFSSL_MLKEM_NO_DECAPSULATE */ #else @@ -7018,9 +7199,10 @@ void mlkem_to_msg(byte* msg, sword16* p) * @param [out] p Polynomial. * @param [in] msg Message as a byte array. */ -void mlkem_from_msg(sword16* p, const byte* msg) +int mlkem_from_msg(sword16* p, const byte* msg) { mlkem_from_msg_neon(p, msg); + return 0; } #endif /* !WOLFSSL_MLKEM_NO_ENCAPSULATE || !WOLFSSL_MLKEM_NO_DECAPSULATE */ @@ -7032,9 +7214,10 @@ void mlkem_from_msg(sword16* p, const byte* msg) * @param [out] msg Message as a byte array. * @param [in, out] p Polynomial. */ -void mlkem_to_msg(byte* msg, sword16* p) +int mlkem_to_msg(byte* msg, sword16* p) { mlkem_to_msg_neon(msg, p); + return 0; } #endif /* WOLFSSL_MLKEM_NO_DECAPSULATE */ #endif /* !(__aarch64__ && WOLFSSL_ARMASM) */ @@ -7080,15 +7263,17 @@ static void mlkem_from_bytes_c(sword16* p, const byte* b, int k) * @param [in] b Array of bytes. * @param [in] k Number of polynomials in vector. */ -void mlkem_from_bytes(sword16* p, const byte* b, int k) +int mlkem_from_bytes(sword16* p, const byte* b, int k) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < k; i++) { mlkem_from_bytes_avx512_vbmi(p, b); p += MLKEM_N; @@ -7100,9 +7285,12 @@ void mlkem_from_bytes(sword16* p, const byte* b, int k) else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < k; i++) { mlkem_from_bytes_avx512(p, b); p += MLKEM_N; @@ -7113,9 +7301,12 @@ void mlkem_from_bytes(sword16* p, const byte* b, int k) } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < k; i++) { mlkem_from_bytes_avx2(p, b); p += MLKEM_N; @@ -7129,6 +7320,7 @@ void mlkem_from_bytes(sword16* p, const byte* b, int k) { mlkem_from_bytes_c(p, b, k); } + return 0; } /* Convert polynomial to bytes. @@ -7175,15 +7367,17 @@ static void mlkem_to_bytes_c(byte* b, sword16* p, int k) * @param [in, out] p Polynomial. * @param [in] k Number of polynomials in vector. */ -void mlkem_to_bytes(byte* b, sword16* p, int k) +int mlkem_to_bytes(byte* b, sword16* p, int k) { #ifdef USE_INTEL_SPEEDUP #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512_VBMI if (USE_INTEL_AVX512(cpuid_flags) && - IS_INTEL_AVX512_VBMI(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX512_VBMI(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < k; i++) { mlkem_to_bytes_avx512_vbmi(b, p); p += MLKEM_N; @@ -7195,9 +7389,12 @@ void mlkem_to_bytes(byte* b, sword16* p, int k) else #endif #ifdef WOLFSSL_MLKEM_HAVE_INTEL_AVX512 - if (USE_INTEL_AVX512(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (USE_INTEL_AVX512(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < k; i++) { mlkem_to_bytes_avx512(b, p); p += MLKEM_N; @@ -7208,9 +7405,12 @@ void mlkem_to_bytes(byte* b, sword16* p, int k) } else #endif - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + if (IS_INTEL_AVX2(cpuid_flags)) { int i; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; for (i = 0; i < k; i++) { mlkem_to_bytes_avx2(b, p); p += MLKEM_N; @@ -7224,6 +7424,7 @@ void mlkem_to_bytes(byte* b, sword16* p, int k) { mlkem_to_bytes_c(b, p, k); } + return 0; } /** diff --git a/wolfcrypt/src/wc_slhdsa.c b/wolfcrypt/src/wc_slhdsa.c index b84566db954..1040d4daac0 100644 --- a/wolfcrypt/src/wc_slhdsa.c +++ b/wolfcrypt/src/wc_slhdsa.c @@ -550,7 +550,14 @@ static int slhdsakey_hash_shake_3(wc_Shake* shake, const byte* data1, #ifndef WC_SHA3_NO_ASM /* Check availability of AVX2 instructions. */ - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + /* CPUID picks the lane; a refused save is an error, never another lane. */ + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + /* The state still holds the seed in the clear, so wipe it. */ + ForceZero(state, sizeof(shake->s)); + return svr_ret; + } /* Process the state using AVX2 instructions. */ sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); @@ -662,7 +669,14 @@ static int slhdsakey_hash_shake_4(wc_Shake* shake, const byte* data1, #ifndef WC_SHA3_NO_ASM /* Check availability of AVX2 instructions. */ - if (IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { + /* CPUID picks the lane; a refused save is an error, never another lane. */ + if (IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) { + /* The state still holds the seed in the clear, so wipe it. */ + ForceZero(state, sizeof(shake->s)); + return svr_ret; + } /* Process the state using AVX2 instructions. */ sha3_block_avx2(state); RESTORE_VECTOR_REGISTERS(); @@ -4719,11 +4733,14 @@ static int slhdsakey_wots_pkgen(SlhDsaKey* key, const byte* sk_seed, #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) && \ !defined(WOLFSSL_SLHDSA_NO_SHAKE) if (!SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = slhdsakey_wots_pkgen_chain_x4(key, sk_seed, pk_seed, adrs, - sk_adrs); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + /* A refused save falls through so the hash below is still freed. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + ret = slhdsakey_wots_pkgen_chain_x4(key, sk_seed, pk_seed, + adrs, sk_adrs); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -4733,11 +4750,14 @@ static int slhdsakey_wots_pkgen(SlhDsaKey* key, const byte* sk_seed, !defined(WOLFSSL_WC_SLHDSA_SMALL) /* The SHA-2 sets batch sixteen chains of SHA-256. */ if (SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = slhdsakey_wots_pkgen_chain_sha2_x16(key, sk_seed, pk_seed, - adrs, sk_adrs); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX512(cpuid_flags)) { + /* CPUID picks the lane; a refused save is an error, not the C lane. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + ret = slhdsakey_wots_pkgen_chain_sha2_x16(key, sk_seed, + pk_seed, adrs, sk_adrs); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -5178,8 +5198,10 @@ static int slhdsakey_wots_sign(SlhDsaKey* key, const byte* m, !defined(WOLFSSL_SLHDSA_NO_SHAKE) /* Steps 11-17: Generate signature from msg. */ if (!SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = slhdsakey_wots_sign_chain_x4(key, msg, sk_seed, pk_seed, adrs, sk_adrs, sig); RESTORE_VECTOR_REGISTERS(); @@ -5192,11 +5214,14 @@ static int slhdsakey_wots_sign(SlhDsaKey* key, const byte* m, !defined(WOLFSSL_WC_SLHDSA_SMALL) /* The SHA-2 sets batch sixteen chains of SHA-256. */ if (SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = slhdsakey_wots_sign_chain_sha2_x16(key, msg, sk_seed, - pk_seed, adrs, sk_adrs, sig); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX512(cpuid_flags)) { + /* CPUID picks the lane; a refused save is an error, not the C lane. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + ret = slhdsakey_wots_sign_chain_sha2_x16(key, msg, sk_seed, + pk_seed, adrs, sk_adrs, sig); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -5845,8 +5870,10 @@ static int slhdsakey_wots_pk_from_sig(SlhDsaKey* key, const byte* sig, #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) && \ !defined(WOLFSSL_SLHDSA_NO_SHAKE) if (!SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { + int svr_ret = SAVE_VECTOR_REGISTERS2(); + if (svr_ret != 0) + return svr_ret; ret = slhdsakey_wots_pk_from_sig_x4(key, sig, msg, pk_seed, adrs, pk_sig); RESTORE_VECTOR_REGISTERS(); @@ -5859,11 +5886,14 @@ static int slhdsakey_wots_pk_from_sig(SlhDsaKey* key, const byte* sig, !defined(WOLFSSL_WC_SLHDSA_SMALL) /* The SHA-2 sets batch sixteen chains of SHA-256. */ if (SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX512(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = slhdsakey_wots_pk_from_sig_sha2_x16(key, sig, msg, - pk_seed, adrs, pk_sig); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX512(cpuid_flags)) { + /* CPUID picks the lane; a refused save is an error, not the C lane. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + ret = slhdsakey_wots_pk_from_sig_sha2_x16(key, sig, msg, + pk_seed, adrs, pk_sig); + RESTORE_VECTOR_REGISTERS(); + } } else #endif @@ -7552,6 +7582,7 @@ static int slhdsakey_fors_sign(SlhDsaKey* key, const byte* md, const byte* sk_seed, const byte* pk_seed, word32* adrs, byte* sig_fors) { int ret = WC_NO_ERR_TRACE(BAD_FUNC_ARG); + byte* sig_start = sig_fors; word16 indices[SLHDSA_MAX_INDICES_SZ]; int i; int j; @@ -7568,6 +7599,8 @@ static int slhdsakey_fors_sign(SlhDsaKey* key, const byte* md, ret = slhdsakey_fors_sk_gen(key, sk_seed, pk_seed, adrs, ((word32)i << a) + indices[i], sig_fors); if (ret != 0) { + /* This slot may already hold a private key value. */ + ForceZero(sig_fors, n); break; } /* Step 4: Move over private key value. */ @@ -7576,9 +7609,14 @@ static int slhdsakey_fors_sign(SlhDsaKey* key, const byte* md, #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) && \ !defined(WOLFSSL_SLHDSA_NO_SHAKE) if (!SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { + IS_INTEL_AVX2(cpuid_flags)) { word16 idx = indices[i]; + int svr_ret = SAVE_VECTOR_REGISTERS2(); + + if (svr_ret != 0) { + ret = svr_ret; + break; + } /* Step 5: For each bit: */ for (j = 0; j < a; j++) { /* Calculate side. */ @@ -7588,6 +7626,8 @@ static int slhdsakey_fors_sign(SlhDsaKey* key, const byte* md, ((word32)i << (a - j)) + s, (word32)j, pk_seed, adrs, sig_fors); if (ret != 0) { + /* At j == 0 this slot holds a private key value. */ + ForceZero(sig_fors, n); break; } /* Step 9: Move signature to after authentication node. */ @@ -7610,6 +7650,8 @@ static int slhdsakey_fors_sign(SlhDsaKey* key, const byte* md, ((word32)i << (a - j)) + s, (word32)j, pk_seed, adrs, sig_fors); if (ret != 0) { + /* At j == 0 this slot holds a private key value. */ + ForceZero(sig_fors, n); break; } /* Step 9: Move signature to after authentication node. */ @@ -7623,6 +7665,11 @@ static int slhdsakey_fors_sign(SlhDsaKey* key, const byte* md, } } + if (ret != 0) { + /* Private key values reached the caller's buffer; do not leave them. */ + ForceZero(sig_start, (size_t)(sig_fors - sig_start)); + } + return ret; } #endif /* !WOLFSSL_SLHDSA_VERIFY_ONLY */ @@ -8482,11 +8529,14 @@ static int slhdsakey_fors_pk_from_sig(SlhDsaKey* key, const byte* sig_fors, #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) && \ !defined(WOLFSSL_SLHDSA_NO_SHAKE) if ((ret == 0) && !SLHDSA_IS_SHA2(key->params->param) && - IS_INTEL_AVX2(cpuid_flags) && - (SAVE_VECTOR_REGISTERS2() == 0)) { - ret = slhdsakey_fors_pk_from_sig_x4(key, sig_fors, indices, pk_seed, - adrs); - RESTORE_VECTOR_REGISTERS(); + IS_INTEL_AVX2(cpuid_flags)) { + /* A refused save falls through so the hash below is still freed. */ + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + ret = slhdsakey_fors_pk_from_sig_x4(key, sig_fors, indices, + pk_seed, adrs); + RESTORE_VECTOR_REGISTERS(); + } } else #endif diff --git a/wolfcrypt/src/wc_xmss_impl.c b/wolfcrypt/src/wc_xmss_impl.c index e6101974427..6d0021427e9 100644 --- a/wolfcrypt/src/wc_xmss_impl.c +++ b/wolfcrypt/src/wc_xmss_impl.c @@ -3255,10 +3255,15 @@ static void wc_xmss_wots_gen_pk(XmssState* state, const byte* sk, /* Every chain here runs the full XMSS_WOTS_W - 1 steps, so the lanes * stay in step and the batch needs no scheduling. */ lanes = XMSS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_xmss_wots_pk_chains_n_way(state, seed, - addr_buf, lanes, pk); - RESTORE_VECTOR_REGISTERS(); + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + int ret = SAVE_VECTOR_REGISTERS2(); + + if (ret == 0) { + ret = wc_xmss_wots_pk_chains_n_way(state, seed, + addr_buf, lanes, pk); + RESTORE_VECTOR_REGISTERS(); + } if (state->ret == 0) { state->ret = ret; } @@ -3287,10 +3292,15 @@ static void wc_xmss_wots_gen_pk(XmssState* state, const byte* sk, /* SHAKE parameter sets of the right shape batch here; everything * else falls through to the chain-at-a-time code below. */ lanes = XMSS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_xmss_wots_pk_chains_n_way(state, seed, - addr_buf, lanes, pk); - RESTORE_VECTOR_REGISTERS(); + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + int ret = SAVE_VECTOR_REGISTERS2(); + + if (ret == 0) { + ret = wc_xmss_wots_pk_chains_n_way(state, seed, + addr_buf, lanes, pk); + RESTORE_VECTOR_REGISTERS(); + } if (state->ret == 0) { state->ret = ret; } @@ -3359,10 +3369,15 @@ static void wc_xmss_wots_sign(XmssState* state, const byte* m, /* Chain i stops at msg[i], so a batch runs as long as its longest * chain and the shorter ones idle at their final value. */ lanes = XMSS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, - NULL, state->encMsg, lanes, sig, sig); - RESTORE_VECTOR_REGISTERS(); + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + int ret = SAVE_VECTOR_REGISTERS2(); + + if (ret == 0) { + ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, + NULL, state->encMsg, lanes, sig, sig); + RESTORE_VECTOR_REGISTERS(); + } if (state->ret == 0) { state->ret = ret; } @@ -3390,10 +3405,15 @@ static void wc_xmss_wots_sign(XmssState* state, const byte* m, #ifdef WC_XMSS_N_WAY /* SHAKE parameter sets of the right shape batch here. */ lanes = XMSS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, - NULL, state->encMsg, lanes, sig, sig); - RESTORE_VECTOR_REGISTERS(); + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + int ret = SAVE_VECTOR_REGISTERS2(); + + if (ret == 0) { + ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, + NULL, state->encMsg, lanes, sig, sig); + RESTORE_VECTOR_REGISTERS(); + } if (state->ret == 0) { state->ret = ret; } @@ -3457,10 +3477,15 @@ static void wc_xmss_wots_pk_from_sig(XmssState* state, const byte* sig, /* Chain i resumes at msg[i] and runs to XMSS_WOTS_W - 1, so a batch * is bounded by the chain of the group that resumed earliest. */ lanes = XMSS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, - state->encMsg, NULL, lanes, sig, pk); - RESTORE_VECTOR_REGISTERS(); + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + int ret = SAVE_VECTOR_REGISTERS2(); + + if (ret == 0) { + ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, + state->encMsg, NULL, lanes, sig, pk); + RESTORE_VECTOR_REGISTERS(); + } if (state->ret == 0) { state->ret = ret; } @@ -3487,10 +3512,15 @@ static void wc_xmss_wots_pk_from_sig(XmssState* state, const byte* sig, #ifdef WC_XMSS_N_WAY /* SHAKE parameter sets of the right shape batch here. */ lanes = XMSS_N_WAY_LANES(params); - if ((lanes > 0) && (SAVE_VECTOR_REGISTERS2() == 0)) { - int ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, - state->encMsg, NULL, lanes, sig, pk); - RESTORE_VECTOR_REGISTERS(); + if (lanes > 0) { + /* CPUID picks the lane; a refused save is an error. */ + int ret = SAVE_VECTOR_REGISTERS2(); + + if (ret == 0) { + ret = wc_xmss_wots_chains_n_way(state, seed, addr_buf, + state->encMsg, NULL, lanes, sig, pk); + RESTORE_VECTOR_REGISTERS(); + } if (state->ret == 0) { state->ret = ret; } diff --git a/wolfssl/wolfcrypt/cpuid.h b/wolfssl/wolfcrypt/cpuid.h index e627dff4699..4d7159d02eb 100644 --- a/wolfssl/wolfcrypt/cpuid.h +++ b/wolfssl/wolfcrypt/cpuid.h @@ -247,7 +247,7 @@ typedef word32 cpuid_flags_t; return 0; } - /* Public APIs to modify flags. */ + /* Public APIs to modify flags. FIPS v7 ignores them, see cpuid.c. */ #ifdef WOLFSSL_API_PREFIX_MAP #define cpuid_select_flags wc_cpuid_select_flags diff --git a/wolfssl/wolfcrypt/settings.h b/wolfssl/wolfcrypt/settings.h index 85f98891574..ac2d2a80f50 100644 --- a/wolfssl/wolfcrypt/settings.h +++ b/wolfssl/wolfcrypt/settings.h @@ -4835,6 +4835,19 @@ #undef DEBUG_VECTOR_REGISTER_ACCESS_FUZZING #endif +/* CPUID pins the lane in these files, so a refused save is an error and + * never a switch to the C code kept for CPUs without the feature. */ +#if (defined(_WC_BUILDING_SP_X86_64_C) || \ + defined(_WC_BUILDING_WC_MLKEM_POLY_C) || \ + defined(_WC_BUILDING_WC_MLDSA_C) || \ + defined(_WC_BUILDING_WC_SLHDSA_C) || \ + defined(_WC_BUILDING_WC_LMS_IMPL_C) || \ + defined(_WC_BUILDING_WC_XMSS_IMPL_C)) && \ + defined(DEBUG_VECTOR_REGISTER_ACCESS_FUZZING) && \ + !defined(DEBUG_FORCE_VECTOR_REGISTER_ACCESS_FUZZING) + #undef DEBUG_VECTOR_REGISTER_ACCESS_FUZZING +#endif + /* Make sure setting OPENSSL_ALL also sets OPENSSL_EXTRA. */ #if defined(OPENSSL_ALL) && !defined(OPENSSL_EXTRA) #define OPENSSL_EXTRA @@ -6177,6 +6190,12 @@ blinding by defining WC_BLINDING_NO_RNG_ACKNOWLEDGE_WEAKNESS." #undef WOLFSSL_CSHAKE #endif +#if defined(WC_C_DYNAMIC_FALLBACK) && defined(HAVE_FIPS) && \ + FIPS_VERSION3_GE(7,0,0) && !defined(WOLFSSL_FIPS_DEV) && \ + !defined(WOLFSSL_FIPS_DEV_NO_POST) + #error WC_C_DYNAMIC_FALLBACK needs --enable-fips=dev or dev-no-post +#endif + /* setup for opt-in DH in FIPS v7+ */ #if FIPS_VERSION3_GE(7,0,0) && !defined(HAVE_DH) && !defined(NO_DH) #define NO_DH diff --git a/wolfssl/wolfcrypt/wc_mldsa.h b/wolfssl/wolfcrypt/wc_mldsa.h index 34ba33efc26..4f3cf8ac80b 100644 --- a/wolfssl/wolfcrypt/wc_mldsa.h +++ b/wolfssl/wolfcrypt/wc_mldsa.h @@ -1103,10 +1103,10 @@ WOLFSSL_API int wc_MlDsaKey_GetSigLen(wc_MlDsaKey* key, int* len); #if !defined(WOLFSSL_MLDSA_NO_SIGN) || \ !defined(WOLFSSL_MLDSA_NO_VERIFY) #ifndef WOLFSSL_NO_ML_DSA_44 -WOLFSSL_TEST_VIS void wc_mldsa_encode_w1_88(const sword32* w1, byte* w1e); +WOLFSSL_TEST_VIS int wc_mldsa_encode_w1_88(const sword32* w1, byte* w1e); #endif #if !defined(WOLFSSL_NO_ML_DSA_65) || !defined(WOLFSSL_NO_ML_DSA_87) -WOLFSSL_TEST_VIS void wc_mldsa_encode_w1_32(const sword32* w1, byte* w1e); +WOLFSSL_TEST_VIS int wc_mldsa_encode_w1_32(const sword32* w1, byte* w1e); #endif #endif diff --git a/wolfssl/wolfcrypt/wc_mlkem.h b/wolfssl/wolfcrypt/wc_mlkem.h index b1ba69a4100..307cd5d65c8 100644 --- a/wolfssl/wolfcrypt/wc_mlkem.h +++ b/wolfssl/wolfcrypt/wc_mlkem.h @@ -600,35 +600,35 @@ WOLFSSL_LOCAL void mlkem_prf_free(MLKEM_PRF_T* prf); WOLFSSL_LOCAL -int mlkem_cmp(const byte* a, const byte* b, int sz); +int mlkem_cmp(const byte* a, const byte* b, int sz, int* fail); WOLFSSL_LOCAL -void mlkem_vec_compress_10(byte* r, sword16* v, unsigned int kp); +int mlkem_vec_compress_10(byte* r, sword16* v, unsigned int kp); WOLFSSL_LOCAL -void mlkem_vec_compress_11(byte* r, sword16* v); +int mlkem_vec_compress_11(byte* r, sword16* v); WOLFSSL_LOCAL -void mlkem_vec_decompress_10(sword16* v, const unsigned char* b, +int mlkem_vec_decompress_10(sword16* v, const unsigned char* b, unsigned int kp); WOLFSSL_LOCAL -void mlkem_vec_decompress_11(sword16* v, const unsigned char* b); +int mlkem_vec_decompress_11(sword16* v, const unsigned char* b); WOLFSSL_LOCAL -void mlkem_compress_4(byte* b, sword16* p); +int mlkem_compress_4(byte* b, sword16* p); WOLFSSL_LOCAL -void mlkem_compress_5(byte* b, sword16* p); +int mlkem_compress_5(byte* b, sword16* p); WOLFSSL_LOCAL -void mlkem_decompress_4(sword16* p, const unsigned char* b); +int mlkem_decompress_4(sword16* p, const unsigned char* b); WOLFSSL_LOCAL -void mlkem_decompress_5(sword16* p, const unsigned char* b); +int mlkem_decompress_5(sword16* p, const unsigned char* b); WOLFSSL_LOCAL -void mlkem_from_msg(sword16* p, const byte* msg); +int mlkem_from_msg(sword16* p, const byte* msg); WOLFSSL_LOCAL -void mlkem_to_msg(byte* msg, sword16* p); +int mlkem_to_msg(byte* msg, sword16* p); WOLFSSL_LOCAL -void mlkem_from_bytes(sword16* p, const byte* b, int k); +int mlkem_from_bytes(sword16* p, const byte* b, int k); WOLFSSL_LOCAL -void mlkem_to_bytes(byte* b, sword16* p, int k); +int mlkem_to_bytes(byte* b, sword16* p, int k); WOLFSSL_LOCAL int mlkem_check_reduced(const sword16* p, int k);