From c84798824fac1a16b83de3e5c850c223bd1e80eb Mon Sep 17 00:00:00 2001 From: Tobias Frauenschlaeger Date: Tue, 18 Aug 2026 09:02:58 +0000 Subject: [PATCH 1/6] Report availability of wc_Sha256HashBlock with a macro The prototype in sha256.h was guarded by WOLFSSL_HAVE_LMS together with WOLFSSL_LMS_FULL_HASH, and nothing in the tree ever defines the latter. The definition, meanwhile, is only compiled in the arm that selects a software transform. On a port that supplies its own Update and Final, LMS therefore declared and called a function that was never built. Add WOLFSSL_HAVE_SHA256_HASH_BLOCK to sha256.h, listing in one place the ports where the helper is absent, and use it for both the prototype and the definition. The list is negative so that naming a port that does have the transform only costs the fast path, while a build combining two ports falls back instead of guessing. A compile time check in sha256.c catches the list drifting away from the arms that select XTRANSFORM. Move every in-tree caller onto the new macro in the same commit: the wolfCrypt and API test suites and the MC/DC hash fault injector all gated on WOLFSSL_LMS_FULL_HASH too, so a build without raw hash access kept calling a function that was not compiled. WOLFSSL_LMS_FULL_HASH is now referenced nowhere and leaves the known macro list. LMS now falls back to the full hash API when the helper is missing. --- .wolfssl_known_macro_extras | 1 - tests/api/test_sha256.c | 4 +--- tests/unit-mcdc/mcdc_fault_hash.h | 4 ++-- wolfcrypt/src/sha256.c | 18 ++++++++++++++++-- wolfcrypt/test/test.c | 8 +++----- wolfssl/wolfcrypt/sha256.h | 31 ++++++++++++++++++++++++++++++- wolfssl/wolfcrypt/wc_lms.h | 4 ++++ 7 files changed, 56 insertions(+), 14 deletions(-) diff --git a/.wolfssl_known_macro_extras b/.wolfssl_known_macro_extras index 4f53192b6d2..d97ae026696 100644 --- a/.wolfssl_known_macro_extras +++ b/.wolfssl_known_macro_extras @@ -987,7 +987,6 @@ WOLFSSL_LINUXKM_USE_GET_RANDOM_KPROBES WOLFSSL_LINUXKM_USE_GET_RANDOM_USER_KRETPROBE WOLFSSL_LINUXKM_USE_MUTEXES WOLFSSL_LMS_CACHE_BITS -WOLFSSL_LMS_FULL_HASH WOLFSSL_LMS_MAX_HEIGHT WOLFSSL_LMS_MAX_LEVELS WOLFSSL_LMS_ROOT_LEVELS diff --git a/tests/api/test_sha256.c b/tests/api/test_sha256.c index d6bdec5a722..bba029721d6 100644 --- a/tests/api/test_sha256.c +++ b/tests/api/test_sha256.c @@ -220,9 +220,7 @@ int test_wc_Sha256Transform(void) int test_wc_Sha256HashBlock_unaligned(void) { EXPECT_DECLS; -#if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) && \ - !defined(NO_SHA256) && !defined(WOLFSSL_KCAPI_HASH) && \ - !defined(WOLF_CRYPTO_CB_ONLY_SHA256) +#ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK wc_Sha256 sha256; byte buf[WC_SHA256_BLOCK_SIZE * 2]; byte aligned[WC_SHA256_DIGEST_SIZE]; diff --git a/tests/unit-mcdc/mcdc_fault_hash.h b/tests/unit-mcdc/mcdc_fault_hash.h index 5ef7708c56c..ccb66e747e8 100644 --- a/tests/unit-mcdc/mcdc_fault_hash.h +++ b/tests/unit-mcdc/mcdc_fault_hash.h @@ -207,7 +207,7 @@ MCDC_FH_MAYBE_UNUSED static int mcdc_fh_Sha256Final(wc_Sha256* sha, byte* hash) return MCDC_FH_ERR; return wc_Sha256Final(sha, hash); } -#if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) +#ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK MCDC_FH_MAYBE_UNUSED static int mcdc_fh_Sha256HashBlock(wc_Sha256* sha, const unsigned char* data, unsigned char* hash) { @@ -385,7 +385,7 @@ MCDC_FH_MAYBE_UNUSED static int mcdc_fh_AesSetKeyDirect(Aes* aes, const byte* ke #define wc_Sha256Update(a, b, c) mcdc_fh_Sha256Update((a), (b), (c)) #undef wc_Sha256Final #define wc_Sha256Final(a, b) mcdc_fh_Sha256Final((a), (b)) - #if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) + #ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK #undef wc_Sha256HashBlock #define wc_Sha256HashBlock(a, b, c) \ mcdc_fh_Sha256HashBlock((a), (b), (c)) diff --git a/wolfcrypt/src/sha256.c b/wolfcrypt/src/sha256.c index 05ceeea874f..597d6b9ea8e 100644 --- a/wolfcrypt/src/sha256.c +++ b/wolfcrypt/src/sha256.c @@ -146,6 +146,14 @@ on the specific device platform. #endif #endif +/* The arms below and the WOLFSSL_KCAPI_HASH block compile no software + * transform, so the cross-check inside the software arm never reaches them. */ +#if defined(WOLFSSL_HAVE_SHA256_HASH_BLOCK) && \ + (defined(WOLFSSL_TI_HASH) || defined(WOLFSSL_CRYPTOCELL) || \ + defined(MAX3266X_SHA) || defined(WOLFSSL_KCAPI_HASH)) + #error "WOLFSSL_HAVE_SHA256_HASH_BLOCK set without a SHA-256 transform" +#endif + #if defined(WOLFSSL_TI_HASH) /* #include included by wc_port.c */ #elif defined(WOLFSSL_CRYPTOCELL) @@ -1884,6 +1892,12 @@ static WC_INLINE int Transform_Sha256_Len(wc_Sha256* sha256, const byte* data, #endif /* End wc_ software implementation */ +/* The port list behind WOLFSSL_HAVE_SHA256_HASH_BLOCK in sha256.h has to + * agree with the arms that select a transform here. */ +#if defined(WOLFSSL_HAVE_SHA256_HASH_BLOCK) && !defined(XTRANSFORM) + #error "WOLFSSL_HAVE_SHA256_HASH_BLOCK set without a SHA-256 transform" +#endif + #ifdef XTRANSFORM static WC_INLINE void AddLength(wc_Sha256* sha256, word32 len) @@ -2422,7 +2436,7 @@ static WC_INLINE int Transform_Sha256_Len(wc_Sha256* sha256, const byte* data, } #endif /* OPENSSL_EXTRA || HAVE_CURL */ -#if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) +#ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK /* One block will be used from data. * hash must be big enough to hold all of digest output. */ @@ -2531,7 +2545,7 @@ static WC_INLINE int Transform_Sha256_Len(wc_Sha256* sha256, const byte* data, return ret; } -#endif /* WOLFSSL_HAVE_LMS && !WOLFSSL_LMS_FULL_HASH */ +#endif /* WOLFSSL_HAVE_SHA256_HASH_BLOCK */ #endif /* !WOLFSSL_KCAPI_HASH */ #endif /* XTRANSFORM */ diff --git a/wolfcrypt/test/test.c b/wolfcrypt/test/test.c index 09cd3d50e30..8e32a244ca2 100644 --- a/wolfcrypt/test/test.c +++ b/wolfcrypt/test/test.c @@ -6521,8 +6521,7 @@ static wc_test_ret_t sha256_large_hash_test(wc_Sha256* sha) #undef LARGE_HASH_TEST_INPUT_SZ #endif /* NO_LARGE_HASH_TEST */ -#if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) && \ - !defined(WOLF_CRYPTO_CB_ONLY_SHA256) +#ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK static wc_test_ret_t sha256_lms_test(wc_Sha256* sha) { byte hash[WC_SHA256_DIGEST_SIZE]; @@ -6559,7 +6558,7 @@ static wc_test_ret_t sha256_lms_test(wc_Sha256* sha) wc_Sha256Free(sha); return ret; } -#endif /* WOLFSSL_HAVE_LMS && !WOLFSSL_LMS_FULL_HASH */ +#endif /* WOLFSSL_HAVE_SHA256_HASH_BLOCK */ #if !defined(HAVE_SELFTEST) && (!defined(HAVE_FIPS) || FIPS_VERSION_GE(7, 0)) static wc_test_ret_t sha256_copy_test(wc_Sha256* sha, wc_Sha256* shaCopy) @@ -6615,8 +6614,7 @@ WOLFSSL_TEST_SUBROUTINE wc_test_ret_t sha256_test(void) if ((ret = sha256_large_hash_test(&sha)) != 0) return ret; #endif -#if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) && \ - !defined(WOLF_CRYPTO_CB_ONLY_SHA256) +#ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK if ((ret = sha256_lms_test(&sha)) != 0) return ret; #endif diff --git a/wolfssl/wolfcrypt/sha256.h b/wolfssl/wolfcrypt/sha256.h index 4fab6d1de8d..586ed7b14de 100644 --- a/wolfssl/wolfcrypt/sha256.h +++ b/wolfssl/wolfcrypt/sha256.h @@ -112,6 +112,35 @@ #define WOLFSSL_NO_HASH_RAW #endif +/* Ports that replace Update and Final with hardware calls do not build + * wc_Sha256HashBlock(). The list is negative on purpose: naming a port that + * has it only costs the fast path, missing one is a link error. */ +#if (defined(WOLFSSL_HAVE_LMS) || defined(WOLFSSL_HAVE_SLHDSA)) && \ + !defined(WOLFSSL_NO_HASH_RAW) && \ + !defined(WOLFSSL_TI_HASH) && \ + !defined(WOLFSSL_CRYPTOCELL) && \ + !defined(MAX3266X_SHA) && \ + !defined(FREESCALE_LTC_SHA) && \ + !defined(WOLFSSL_PIC32MZ_HASH) && \ + !defined(STM32_HASH_SHA2) && \ + !(defined(WOLFSSL_IMX6_CAAM) && !defined(NO_IMX6_CAAM_HASH)) && \ + !(defined(WOLFSSL_SE050) && defined(WOLFSSL_SE050_HASH)) && \ + !defined(WOLFSSL_AFALG_HASH) && \ + !defined(WOLFSSL_DEVCRYPTO_HASH) && \ + !defined(WOLFSSL_USE_ESP32_CRYPT_HASH_HW) && \ + !defined(WOLFSSL_RENESAS_TSIP_TLS) && \ + !defined(WOLFSSL_RENESAS_TSIP_CRYPTONLY) && \ + !defined(WOLFSSL_RENESAS_SCEPROTECT) && \ + !defined(WOLFSSL_RENESAS_RSIP) && \ + !defined(WOLFSSL_RENESAS_RX64_HASH) && \ + !defined(PSOC6_HASH_SHA2) && \ + !defined(WOLFSSL_IMXRT_DCP) && \ + !defined(WOLFSSL_NXP_HASHCRYPT_SHA) && \ + !defined(WOLFSSL_SILABS_SE_ACCEL) && \ + !defined(WOLFSSL_KCAPI_HASH) + #define WOLFSSL_HAVE_SHA256_HASH_BLOCK +#endif + #define SHA256_NOINLINE WC_NO_INLINE #if !defined(NO_OLD_SHA_NAMES) @@ -274,7 +303,7 @@ WOLFSSL_API int wc_Sha256Reset(wc_Sha256* sha256); !defined(WOLF_CRYPTO_CB_ONLY_SHA256) WOLFSSL_API int wc_Sha256Transform(wc_Sha256* sha, const unsigned char* data); #endif -#if defined(WOLFSSL_HAVE_LMS) && !defined(WOLFSSL_LMS_FULL_HASH) +#ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK WOLFSSL_API int wc_Sha256HashBlock(wc_Sha256* sha, const unsigned char* data, unsigned char* hash); #endif diff --git a/wolfssl/wolfcrypt/wc_lms.h b/wolfssl/wolfcrypt/wc_lms.h index ed3ea34a19e..5de55141dc1 100644 --- a/wolfssl/wolfcrypt/wc_lms.h +++ b/wolfssl/wolfcrypt/wc_lms.h @@ -107,6 +107,10 @@ #if defined(WOLFSSL_NO_HASH_RAW) && !defined(WC_LMS_FULL_HASH) #define WC_LMS_FULL_HASH #endif +/* The SHA-256 parameter sets need the one block compression helper. */ +#if !defined(WOLFSSL_HAVE_SHA256_HASH_BLOCK) && !defined(WC_LMS_FULL_HASH) + #define WC_LMS_FULL_HASH +#endif /* Length of the Key ID. */ #define WC_LMS_I_LEN 16 From 107e6c54e06b00a17365e3e776ade23a7a293326 Mon Sep 17 00:00:00 2001 From: Tobias Frauenschlaeger Date: Tue, 18 Aug 2026 09:03:12 +0000 Subject: [PATCH 2/6] Speed up SLH-DSA hashing Two independent changes to the hot path. The SHA2 parameter sets ran every F, H and PRF through the streaming hash API: a free, a context copy, two updates and a final, all to compute one 64 byte compression. The free was redundant, because wc_Sha256Copy releases the destination itself, and the rest can be replaced by building the block that follows the pre-computed PK.seed midstate and compressing it once. The four helpers now share one dispatcher, which takes wc_Sha256HashBlock where it is available and the streaming API otherwise. Measured on a Xeon at 2.1 GHz, SHA2-128f signing goes from 28.1 ms to 22.7 ms. The SHAKE parameter sets spend about 93 percent of a signature in the four-way AVX2 Keccak permutation, while the eight-way AVX512 permutation that ML-DSA, ML-KEM and FrodoKEM already use sat unused. Measured per lane, eight ways beat four by two to one or better. Add an eight-way path for the WOTS+ chains, dispatched at run time on capable CPUs. SHAKE-128s signing goes from 309 ms to 194 ms, SHAKE-256s from 481 ms to 332 ms and SHAKE-128f from 15.6 ms to 11.2 ms. FORS and the verify side still use the four-way path. The eight-way helpers loop over the lanes rather than unrolling like their four-way counterparts. Setting up the state is a fraction of a percent of a signature, so the unrolling buys nothing worth the risk of an index error. Public keys and signatures are unchanged. Checked byte for byte against the previous code for all twelve parameter sets, with the eight-way path active and with WOLFSSL_SLHDSA_FULL_HASH forcing the streaming API. The PRF output that now takes the wc_Sha256HashBlock path is secret, and on x86-64 CPUs without MOVBE that function byte-reverses the digest through a stack buffer it never cleared. Wipe it. --- .wolfssl_known_macro_extras | 1 + ChangeLog.md | 9 + wolfcrypt/src/sha256.c | 1 + wolfcrypt/src/wc_slhdsa.c | 737 ++++++++++++++++++++++++++++-------- 4 files changed, 583 insertions(+), 165 deletions(-) diff --git a/.wolfssl_known_macro_extras b/.wolfssl_known_macro_extras index d97ae026696..f41c1311e57 100644 --- a/.wolfssl_known_macro_extras +++ b/.wolfssl_known_macro_extras @@ -1129,6 +1129,7 @@ WOLFSSL_SHA3_PPC64_BLOCKS_N WOLFSSL_SHUTDOWNONCE WOLFSSL_SILABS_TRNG WOLFSSL_SLHDSA_FULL_HASH +WOLFSSL_SLHDSA_NO_SHAKE_X8 WOLFSSL_SLHDSA_NO_VERIFY_ONLY WOLFSSL_SNIFFER_NO_RECOVERY WOLFSSL_SP_FAST_NCT_EXPTMOD diff --git a/ChangeLog.md b/ChangeLog.md index 2c2b6e95a2b..fc1599cd74c 100644 --- a/ChangeLog.md +++ b/ChangeLog.md @@ -622,6 +622,11 @@ PR stands for Pull Request, and PR references a GitHub pull request num * Migrate internal ML-KEM consumers to canonical wc_MlKemKey API by @Frauschi (PR 10571) * Add PQ documentation for LMS, ML-DSA, ML-KEM, XMSS by @kaleb-himes (PR 10514) * Various leak / alloc and zeroization fixes for SLH-DSA by @Frauschi (PR 10698) +* Hash the block after the pre-computed PK.seed midstate directly for the + SLH-DSA SHA2 parameter sets, rather than copying a hash object and + streaming two updates into it for every hash. by @Frauschi +* Use the 8-way AVX512 Keccak permutation for the SLH-DSA WOTS+ chains on + capable CPUs, which the four-way AVX2 path had to itself. by @Frauschi ## TLS/DTLS @@ -729,6 +734,10 @@ PR stands for Pull Request, and PR references a GitHub pull request num * Fixed the SGX build to not require fcntl.h. by @JacobBarthelmeh (PR 10524) * Various linuxkm Fenrir fixes by @douzzer (PR 10688) * Various bsdkm fixes and cleanup by @philljj (PR 10565) +* Added `WOLFSSL_HAVE_SHA256_HASH_BLOCK` to report whether + `wc_Sha256HashBlock()` is built. LMS gated its use on a macro that was + never defined, so ports without the software SHA-256 transform called a + function that was not compiled. by @Frauschi ## Bug Fixes diff --git a/wolfcrypt/src/sha256.c b/wolfcrypt/src/sha256.c index 597d6b9ea8e..fc292972c2d 100644 --- a/wolfcrypt/src/sha256.c +++ b/wolfcrypt/src/sha256.c @@ -2479,6 +2479,7 @@ static WC_INLINE int Transform_Sha256_Len(wc_Sha256* sha256, const byte* data, word32 buf[WC_SHA256_DIGEST_SIZE / sizeof(word32)]; ByteReverseWords(buf, sha256->digest, WC_SHA256_DIGEST_SIZE); XMEMCPY(hash, buf, WC_SHA256_DIGEST_SIZE); + ForceZero(buf, sizeof(buf)); } #endif else { diff --git a/wolfcrypt/src/wc_slhdsa.c b/wolfcrypt/src/wc_slhdsa.c index c2f34ebbdb1..b6dc3bcd444 100644 --- a/wolfcrypt/src/wc_slhdsa.c +++ b/wolfcrypt/src/wc_slhdsa.c @@ -173,6 +173,22 @@ wc_static_assert(SLHDSA_MAX_MSG_SZ <= 255); * declarations and the ForceZero() sizes so they cannot drift. */ #define SLHDSA_SHAKE_X4_STATE_W (25 * 4) +/* Number of word64 in the 8-way (x8) Keccak state used by the AVX512 hash + * helpers (25 lanes * 8 ways). */ +#define SLHDSA_SHAKE_X8_STATE_W (25 * 8) + +/* Words of seed and encoded HashAddress held across an 8-way chain: at most + * four words of seed and four of address, one copy per lane. */ +#define SLHDSA_SHAKE_X8_FIXED_W (8 * 8) + +/* The 8-way AVX512 Keccak permutation is built with the rest of the SHA-3 + * assembly, so the WOTS+ chains dispatch to it on capable CPUs. */ +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) && \ + !defined(NO_AVX512_SUPPORT) && !defined(WOLFSSL_SLHDSA_NO_SHAKE) && \ + !defined(WOLFSSL_SLHDSA_NO_SHAKE_X8) + #define SLHDSA_HAVE_SHAKE_X8 +#endif + #ifndef WC_SLHDSA_ALL_NO_256F /* Maximum number of bytes to produce from digest of message. */ #define SLHDSA_MAX_MD 49 @@ -695,6 +711,26 @@ static int slhdsakey_hash_shake_4(wc_Shake* shake, const byte* data1, /* Size of compressed HashAddress (ADRS^c) per FIPS 205 Section 11.2. */ #define SLHDSA_HAC_SZ 22 +/* Compress the block after the PK.seed midstate directly rather than copy a + * hash object per hash. WOLF_CRYPTO_CB_FIND consults a callback for every hash + * object, leaving no object the direct path may touch. */ +#if !defined(WOLFSSL_SLHDSA_FULL_HASH) && \ + defined(WOLFSSL_HAVE_SHA256_HASH_BLOCK) && \ + !(defined(WOLF_CRYPTO_CB) && defined(WOLF_CRYPTO_CB_FIND)) + #define SLHDSA_SHA2_BLOCK_HASH +#endif + +#ifdef SLHDSA_SHA2_BLOCK_HASH +/* A registered callback expects to see every update, so the direct path is + * only taken for an object no callback has claimed. */ +#ifdef WOLF_CRYPTO_CB + #define SLHDSA_SHA256_RAW_OK(key) \ + ((key)->hash.sha2.sha256.devId == INVALID_DEVID) +#else + #define SLHDSA_SHA256_RAW_OK(key) 1 +#endif +#endif /* SLHDSA_SHA2_BLOCK_HASH */ + /* Encode a compressed HashAddress (ADRS^c). * * FIPS 205. Section 11.2. @@ -773,59 +809,226 @@ static int slhdsakey_precompute_sha2_midstates(SlhDsaKey* key) return ret; } -/* SHA2 F function. - * - * FIPS 205. Section 11.2. - * F(PK.seed, ADRS, M1) = Trunc_n(SHA-256(PK.seed||pad(64-n)||ADRS^c||M1)) +/* Largest message is n bytes for F and PRF, 2n for H at category 1; both + * leave room for the padding and length in the block after the midstate. */ +wc_static_assert(SLHDSA_HAC_SZ + 32 + 1 + 8 <= WC_SHA256_BLOCK_SIZE); + +#ifdef SLHDSA_SHA2_BLOCK_HASH +/* Hash the block following the pre-computed SHA-256 midstate. * - * Uses pre-computed midstate for the first block. + * F, H and PRF are each one block past the midstate, so the compressed + * address, message and padding are built directly and compressed once. * - * @param [in] key SLH-DSA key (SHA2 hash objects + midstate). - * @param [in] pk_seed Public key seed (unused - midstate). - * @param [in] adrs HashAddress. - * @param [in] m Message of n bytes. - * @param [in] n Number of bytes in hash output. - * @param [out] hash Buffer to hold hash output. + * @param [in] key SLH-DSA key. + * @param [in] address Encoded compressed HashAddress. + * @param [in] m1 First message part. + * @param [in] m1_len Length of first message part. + * @param [in] m2 Second message part, may be NULL. + * @param [in] m2_len Length of second message part. + * @param [out] hash Buffer to hold hash output. + * @param [in] hash_len Number of bytes of hash to output. + * @param [in] zeroize Wipe the working buffers when the input is secret. * @return 0 on success. */ -static int slhdsakey_hash_f_sha2(SlhDsaKey* key, const byte* pk_seed, - const word32* adrs, const byte* m, byte n, byte* hash) +static int slhdsakey_sha256_block_hash(SlhDsaKey* key, const byte* address, + const byte* m1, byte m1_len, const byte* m2, byte m2_len, byte* hash, + byte hash_len) { int ret; - byte address[SLHDSA_HAC_SZ]; + /* wc_Sha256HashBlock() copies any other buffer into this one + * before compressing, so build in place as wc_lms_impl.c does. */ + byte* block = (byte*)key->hash.sha2.sha256.buffer; byte digest[WC_SHA256_DIGEST_SIZE]; + word32 len = (word32)SLHDSA_HAC_SZ + m1_len + m2_len; + /* Length covers the midstate block as well as this one. */ + word32 bits = (WC_SHA256_BLOCK_SIZE + len) * 8; - (void)pk_seed; + XMEMCPY(block, address, SLHDSA_HAC_SZ); + XMEMCPY(block + SLHDSA_HAC_SZ, m1, m1_len); + if (m2_len > 0) { + XMEMCPY(block + SLHDSA_HAC_SZ + m1_len, m2, m2_len); + } + /* SHA-256 padding. */ + block[len] = 0x80; + XMEMSET(block + len + 1, 0, WC_SHA256_BLOCK_SIZE - 8 - (len + 1)); + c32toa(0, block + WC_SHA256_BLOCK_SIZE - 8); + c32toa(bits, block + WC_SHA256_BLOCK_SIZE - 4); - /* Encode compressed address. */ - HA_Encode_Compressed(adrs, address); + /* Restore the midstate and compress. */ + XMEMCPY(key->hash.sha2.sha256.digest, key->hash.sha2.sha256_mid.digest, + WC_SHA256_DIGEST_SIZE); + ret = wc_Sha256HashBlock(&key->hash.sha2.sha256, block, digest); + if (ret == 0) { + XMEMCPY(hash, digest, hash_len); + } - /* Restore SHA-256 midstate. */ +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Add("slhdsa sha256 block", block, WC_SHA256_BLOCK_SIZE); + wc_MemZero_Add("slhdsa sha256 digest", digest, sizeof(digest)); +#endif + ForceZero(block, WC_SHA256_BLOCK_SIZE); + ForceZero(digest, sizeof(digest)); +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Check(digest, sizeof(digest)); + wc_MemZero_Check(block, WC_SHA256_BLOCK_SIZE); +#endif - if (key->hash.sha2.sha256_inited) { - wc_Sha256Free(&key->hash.sha2.sha256); - key->hash.sha2.sha256_inited = 0; - } + return ret; +} +#endif /* SLHDSA_SHA2_BLOCK_HASH */ + +/* Hash the compressed address and message with SHA-256 from the midstate. + * + * @param [in] key SLH-DSA key. + * @param [in] address Encoded compressed HashAddress. + * @param [in] m1 First message part. + * @param [in] m1_len Length of first message part. + * @param [in] m2 Second message part, may be NULL. + * @param [in] m2_len Length of second message part. + * @param [out] hash Buffer to hold hash output. + * @param [in] hash_len Number of bytes of hash to output. + * @return 0 on success. + */ +static int slhdsakey_sha256_api_hash(SlhDsaKey* key, const byte* address, + const byte* m1, byte m1_len, const byte* m2, byte m2_len, byte* hash, + byte hash_len) +{ + int ret; + byte digest[WC_SHA256_DIGEST_SIZE]; + + /* Restore the midstate. wc_Sha256Copy() releases the destination. */ ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, &key->hash.sha2.sha256); if (ret == 0) { key->hash.sha2.sha256_inited = 1; - /* Update with compressed ADRS and message. */ ret = wc_Sha256Update(&key->hash.sha2.sha256, address, SLHDSA_HAC_SZ); } if (ret == 0) { - ret = wc_Sha256Update(&key->hash.sha2.sha256, m, n); + ret = wc_Sha256Update(&key->hash.sha2.sha256, m1, m1_len); + } + if ((ret == 0) && (m2_len > 0)) { + ret = wc_Sha256Update(&key->hash.sha2.sha256, m2, m2_len); } if (ret == 0) { ret = wc_Sha256Final(&key->hash.sha2.sha256, digest); } if (ret == 0) { - /* Truncate to n bytes. */ - XMEMCPY(hash, digest, n); + XMEMCPY(hash, digest, hash_len); } +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Add("slhdsa sha256 digest", digest, sizeof(digest)); +#endif + ForceZero(digest, sizeof(digest)); +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Check(digest, sizeof(digest)); +#endif + return ret; } +/* Hash the compressed address and message with SHA-256. + * + * @param [in] key SLH-DSA key. + * @param [in] address Encoded compressed HashAddress. + * @param [in] m1 First message part. + * @param [in] m1_len Length of first message part. + * @param [in] m2 Second message part, may be NULL. + * @param [in] m2_len Length of second message part. + * @param [out] hash Buffer to hold hash output. + * @param [in] hash_len Number of bytes of hash to output. + * @return 0 on success. + */ +static int slhdsakey_sha256_hash(SlhDsaKey* key, const byte* address, + const byte* m1, byte m1_len, const byte* m2, byte m2_len, byte* hash, + byte hash_len) +{ +#ifdef SLHDSA_SHA2_BLOCK_HASH + if (SLHDSA_SHA256_RAW_OK(key)) { + return slhdsakey_sha256_block_hash(key, address, m1, m1_len, m2, + m2_len, hash, hash_len); + } +#endif + return slhdsakey_sha256_api_hash(key, address, m1, m1_len, m2, m2_len, + hash, hash_len); +} + +/* Hash the compressed address and message with SHA-512 from the midstate. + * + * @param [in] key SLH-DSA key. + * @param [in] address Encoded compressed HashAddress. + * @param [in] m1 First message part. + * @param [in] m1_len Length of first message part. + * @param [in] m2 Second message part, may be NULL. + * @param [in] m2_len Length of second message part. + * @param [out] hash Buffer to hold hash output. + * @param [in] hash_len Number of bytes of hash to output. + * @return 0 on success. + */ +static int slhdsakey_sha512_hash(SlhDsaKey* key, const byte* address, + const byte* m1, byte m1_len, const byte* m2, byte m2_len, byte* hash, + byte hash_len) +{ + int ret; + byte digest[WC_SHA512_DIGEST_SIZE]; + + /* Restore the midstate. wc_Sha512Copy() releases the destination. */ + ret = wc_Sha512Copy(&key->hash.sha2.sha512_mid, &key->hash.sha2.sha512); + if (ret == 0) { + key->hash.sha2.sha512_inited = 1; + ret = wc_Sha512Update(&key->hash.sha2.sha512, address, SLHDSA_HAC_SZ); + } + if (ret == 0) { + ret = wc_Sha512Update(&key->hash.sha2.sha512, m1, m1_len); + } + if ((ret == 0) && (m2_len > 0)) { + ret = wc_Sha512Update(&key->hash.sha2.sha512, m2, m2_len); + } + if (ret == 0) { + ret = wc_Sha512Final(&key->hash.sha2.sha512, digest); + } + if (ret == 0) { + XMEMCPY(hash, digest, hash_len); + } + +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Add("slhdsa sha512 digest", digest, sizeof(digest)); +#endif + ForceZero(digest, sizeof(digest)); +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Check(digest, sizeof(digest)); +#endif + + return ret; +} + +/* SHA2 F function. + * + * FIPS 205. Section 11.2. + * F(PK.seed, ADRS, M1) = Trunc_n(SHA-256(PK.seed||pad(64-n)||ADRS^c||M1)) + * + * Uses pre-computed midstate for the first block. + * + * @param [in] key SLH-DSA key (SHA2 hash objects + midstate). + * @param [in] pk_seed Public key seed (unused - midstate). + * @param [in] adrs HashAddress. + * @param [in] m Message of n bytes. + * @param [in] n Number of bytes in hash output. + * @param [out] hash Buffer to hold hash output. + * @return 0 on success. + */ +static int slhdsakey_hash_f_sha2(SlhDsaKey* key, const byte* pk_seed, + const word32* adrs, const byte* m, byte n, byte* hash) +{ + byte address[SLHDSA_HAC_SZ]; + + (void)pk_seed; + + /* Encode compressed address. */ + HA_Encode_Compressed(adrs, address); + + return slhdsakey_sha256_hash(key, address, m, n, NULL, 0, hash, n); +} + #ifndef WOLFSSL_SLHDSA_VERIFY_ONLY /* SHA2 H function. * @@ -854,68 +1057,28 @@ static int slhdsakey_hash_h_sha2(SlhDsaKey* key, const byte* pk_seed, if (n == WC_SLHDSA_N_128) { /* Category 1: use SHA-256. */ - byte digest[WC_SHA256_DIGEST_SIZE]; - - if (key->hash.sha2.sha256_inited) { - wc_Sha256Free(&key->hash.sha2.sha256); - key->hash.sha2.sha256_inited = 0; - } - ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, - &key->hash.sha2.sha256); - if (ret == 0) { - key->hash.sha2.sha256_inited = 1; - ret = wc_Sha256Update(&key->hash.sha2.sha256, address, - SLHDSA_HAC_SZ); - } - if (ret == 0) { - ret = wc_Sha256Update(&key->hash.sha2.sha256, node, 2U * n); - } - if (ret == 0) { - ret = wc_Sha256Final(&key->hash.sha2.sha256, digest); - } - if (ret == 0) { - XMEMCPY(hash, digest, n); - } + ret = slhdsakey_sha256_hash(key, address, node, (byte)(2 * n), NULL, 0, + hash, n); } else { /* Categories 3, 5: use SHA-512. */ - byte digest[WC_SHA512_DIGEST_SIZE]; - - if (key->hash.sha2.sha512_inited) { - wc_Sha512Free(&key->hash.sha2.sha512); - key->hash.sha2.sha512_inited = 0; - } - ret = wc_Sha512Copy(&key->hash.sha2.sha512_mid, - &key->hash.sha2.sha512); - if (ret == 0) { - key->hash.sha2.sha512_inited = 1; - ret = wc_Sha512Update(&key->hash.sha2.sha512, address, - SLHDSA_HAC_SZ); - } - if (ret == 0) { - ret = wc_Sha512Update(&key->hash.sha2.sha512, node, 2U * n); - } - if (ret == 0) { - ret = wc_Sha512Final(&key->hash.sha2.sha512, digest); - } - if (ret == 0) { - XMEMCPY(hash, digest, n); - } + ret = slhdsakey_sha512_hash(key, address, node, (byte)(2 * n), NULL, 0, + hash, n); } return ret; } #endif /* !WOLFSSL_SLHDSA_VERIFY_ONLY */ -/* SHA2 H function with two separate n-byte halves. +/* SHA2 H function over two separate n-byte messages. * - * Same as slhdsakey_hash_h_sha2 but M2 = m1 || m2. + * FIPS 205. Section 11.2. * * @param [in] key SLH-DSA key. * @param [in] pk_seed Public key seed (unused - midstate). * @param [in] adrs HashAddress. - * @param [in] m1 First n bytes of message. - * @param [in] m2 Second n bytes of message. + * @param [in] m1 First message of n bytes. + * @param [in] m2 Second message of n bytes. * @param [in] n Number of bytes in hash output. * @param [out] hash Buffer to hold hash output. * @return 0 on success. @@ -933,59 +1096,11 @@ static int slhdsakey_hash_h_2_sha2(SlhDsaKey* key, const byte* pk_seed, if (n == WC_SLHDSA_N_128) { /* Category 1: use SHA-256. */ - byte digest[WC_SHA256_DIGEST_SIZE]; - - if (key->hash.sha2.sha256_inited) { - wc_Sha256Free(&key->hash.sha2.sha256); - key->hash.sha2.sha256_inited = 0; - } - ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, - &key->hash.sha2.sha256); - if (ret == 0) { - key->hash.sha2.sha256_inited = 1; - ret = wc_Sha256Update(&key->hash.sha2.sha256, address, - SLHDSA_HAC_SZ); - } - if (ret == 0) { - ret = wc_Sha256Update(&key->hash.sha2.sha256, m1, n); - } - if (ret == 0) { - ret = wc_Sha256Update(&key->hash.sha2.sha256, m2, n); - } - if (ret == 0) { - ret = wc_Sha256Final(&key->hash.sha2.sha256, digest); - } - if (ret == 0) { - XMEMCPY(hash, digest, n); - } + ret = slhdsakey_sha256_hash(key, address, m1, n, m2, n, hash, n); } else { /* Categories 3, 5: use SHA-512. */ - byte digest[WC_SHA512_DIGEST_SIZE]; - - if (key->hash.sha2.sha512_inited) { - wc_Sha512Free(&key->hash.sha2.sha512); - key->hash.sha2.sha512_inited = 0; - } - ret = wc_Sha512Copy(&key->hash.sha2.sha512_mid, - &key->hash.sha2.sha512); - if (ret == 0) { - key->hash.sha2.sha512_inited = 1; - ret = wc_Sha512Update(&key->hash.sha2.sha512, address, - SLHDSA_HAC_SZ); - } - if (ret == 0) { - ret = wc_Sha512Update(&key->hash.sha2.sha512, m1, n); - } - if (ret == 0) { - ret = wc_Sha512Update(&key->hash.sha2.sha512, m2, n); - } - if (ret == 0) { - ret = wc_Sha512Final(&key->hash.sha2.sha512, digest); - } - if (ret == 0) { - XMEMCPY(hash, digest, n); - } + ret = slhdsakey_sha512_hash(key, address, m1, n, m2, n, hash, n); } return ret; @@ -1009,46 +1124,14 @@ static int slhdsakey_hash_h_2_sha2(SlhDsaKey* key, const byte* pk_seed, static int slhdsakey_hash_prf_sha2(SlhDsaKey* key, const byte* pk_seed, const byte* sk_seed, const word32* adrs, byte n, byte* hash) { - int ret; byte address[SLHDSA_HAC_SZ]; - byte digest[WC_SHA256_DIGEST_SIZE]; (void)pk_seed; /* Encode compressed address. */ HA_Encode_Compressed(adrs, address); - /* Restore SHA-256 midstate. */ - if (key->hash.sha2.sha256_inited) { - wc_Sha256Free(&key->hash.sha2.sha256); - key->hash.sha2.sha256_inited = 0; - } - ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, &key->hash.sha2.sha256); - if (ret == 0) { - key->hash.sha2.sha256_inited = 1; - ret = wc_Sha256Update(&key->hash.sha2.sha256, address, SLHDSA_HAC_SZ); - } - if (ret == 0) { - ret = wc_Sha256Update(&key->hash.sha2.sha256, sk_seed, n); - } - if (ret == 0) { - ret = wc_Sha256Final(&key->hash.sha2.sha256, digest); - /* digest now holds the secret PRF output (WOTS+/FORS key); register it - * before it is copied out so any later exit is covered. */ -#ifdef WOLFSSL_CHECK_MEM_ZERO - wc_MemZero_Add("slhdsa prf digest", digest, sizeof(digest)); -#endif - } - if (ret == 0) { - XMEMCPY(hash, digest, n); - } - - /* digest holds the secret PRF output (WOTS+/FORS key). */ - ForceZero(digest, sizeof(digest)); -#ifdef WOLFSSL_CHECK_MEM_ZERO - wc_MemZero_Check(digest, sizeof(digest)); -#endif - return ret; + return slhdsakey_sha256_hash(key, address, sk_seed, n, NULL, 0, hash, n); } #endif /* !WOLFSSL_SLHDSA_VERIFY_ONLY */ @@ -1075,10 +1158,6 @@ static int slhdsakey_hash_start_addr_sha2(SlhDsaKey* key, if (n == WC_SLHDSA_N_128) { /* Category 1: SHA-256 -- use sha256_2 (T_l must not collide with * sha256 which is used by F and H). */ - if (key->hash.sha2.sha256_2_inited) { - wc_Sha256Free(&key->hash.sha2.sha256_2); - key->hash.sha2.sha256_2_inited = 0; - } ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, &key->hash.sha2.sha256_2); if (ret == 0) { @@ -1090,10 +1169,6 @@ static int slhdsakey_hash_start_addr_sha2(SlhDsaKey* key, else { /* Categories 3, 5: SHA-512 -- use sha512_2 (T_l must not collide * with sha512 which is used by H). */ - if (key->hash.sha2.sha512_2_inited) { - wc_Sha512Free(&key->hash.sha2.sha512_2); - key->hash.sha2.sha512_2_inited = 0; - } ret = wc_Sha512Copy(&key->hash.sha2.sha512_mid, &key->hash.sha2.sha512_2); if (ret == 0) { @@ -3223,6 +3298,331 @@ static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, } #endif +#ifdef SLHDSA_HAVE_SHAKE_X8 +/* Fill the 8-way state with the seed and encoded HashAddress, one copy per + * lane. + * + * The 8-way helpers loop over the lanes rather than unrolling like the 4-way + * ones. Setting up the state is a fraction of a percent of a signature; the + * Keccak permutation is the rest. + * + * @param [out] state SHAKE-256 x8 state. + * @param [in] seed Seed at the start of each hash. + * @param [in] addr Encoded HashAddress for each hash. + * @param [in] n Number of bytes of seed. + * @return Offset after the seed and HashAddress. + */ +static word32 slhdsakey_shake256_set_seed_ha_x8(word64* state, + const byte* seed, const byte* addr, int n) +{ + int i; + int l; + word32 o = 0; + + for (i = 0; i < n; i += 8) { + word64 v = readUnalignedWord64(seed + i); + + for (l = 0; l < 8; l++) { + state[o + l] = v; + } + o += 8; + } + for (i = 0; i < SLHDSA_HA_SZ; i += 8) { + word64 v = readUnalignedWord64(addr + i); + + for (l = 0; l < 8; l++) { + state[o + l] = v; + } + o += 8; + } + + return o; +} + +/* Append one n-byte hash per lane to the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset to place the hashes at. + * @param [in] hash Eight n-byte hashes. + * @param [in] n Number of bytes in each hash. + */ +static void slhdsakey_shake256_set_hash_x8(word64* state, word32 o, + const byte* hash, int n) +{ + int i; + int l; + + for (i = 0; i < n; i += 8) { + for (l = 0; l < 8; l++) { + state[o + l] = readUnalignedWord64(hash + l * n + i); + } + o += 8; + } +} + +/* Get the eight SHAKE-256 n-byte hash results. + * + * @param [in] state SHAKE-256 x8 state. + * @param [out] hash Buffer to hold eight n-byte hash results. + * @param [in] n Length of each hash in bytes. + */ +static void slhdsakey_shake256_get_hash_x8(const word64* state, byte* hash, + int n) +{ + int i; + int l; + + for (i = 0; i < (n / 8); i++) { + for (l = 0; l < 8; l++) { + writeUnalignedWord64(hash + l * n + i * 8, state[8 * i + l]); + } + } +} + +/* Set the end of the SHAKE-256 x8 state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset to the end of the data. + */ +static void slhdsakey_shake256_set_end_x8(word64* state, word32 o) +{ + int l; + + /* Data end marker. */ + for (l = 0; l < 8; l++) { + state[o + l] = (word64)0x1f; + } + XMEMSET(state + o + 8, 0, + (size_t)(SLHDSA_SHAKE_X8_STATE_W - (o + 8)) * sizeof(word64)); + /* SHAKE-256 state end marker. */ + for (l = 0; l < 8; l++) { + ((word8*)(state + 8 * WC_SHA3_256_COUNT - 8 + l))[7] ^= 0x80; + } +} + +/* Set an incrementing chain address into each lane of the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset of state after the HashAddress. + * @param [in] a Value to set, incrementing for each hash. + */ +static void slhdsakey_shake256_set_chain_addr_x8(word64* state, word32 o, + byte a) +{ + int l; + + for (l = 0; l < 8; l++) { + ((word8*)(state + o - 8 + l))[3] = (word8)(a + l); + } +} + +/* Set the same hash address into each lane of the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset of state after the HashAddress. + * @param [in] a Value to set for each hash. + */ +static void slhdsakey_shake256_set_hash_addr_x8(word64* state, word32 o, + byte a) +{ + int l; + + for (l = 0; l < 8; l++) { + ((word8*)(state + o - 8 + l))[7] = (word8)a; + } +} + +/* PRF eight WOTS+ secret values at consecutive chain addresses. + * + * @param [in] pk_seed Public key seed. + * @param [in] sk_seed Private key seed. + * @param [in] addr Encoded HashAddress. + * @param [in] n Number of bytes in each hash. + * @param [in] ca Chain address start index. + * @param [out] sk Eight n-byte secret values. + * @param [in] heap Dynamic memory allocation hint. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_hash_prf_x8(const byte* pk_seed, const byte* sk_seed, + byte* addr, byte n, byte ca, byte* sk, void* heap) +{ + int ret = 0; + word32 o; + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap); + + (void)heap; + + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + int i; + int l; + + o = slhdsakey_shake256_set_seed_ha_x8(state, pk_seed, addr, n); + slhdsakey_shake256_set_chain_addr_x8(state, o, ca); + /* PRF hashes the private key seed, the same value in every lane. */ + for (i = 0; i < n; i += 8) { + word64 v = readUnalignedWord64(sk_seed + i); + + for (l = 0; l < 8; l++) { + state[o + l] = v; + } + o += 8; + } + slhdsakey_shake256_set_end_x8(state, o); + + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + sha3_blocksx8_avx512(state); + slhdsakey_shake256_get_hash_x8(state, sk, n); + RESTORE_VECTOR_REGISTERS(); + } + + /* state holds the secret PRF output (WOTS+ key). */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + } + + return ret; +} + +/* Iterate the hash function 15 times over eight hashes. + * + * FIPS 205. Section 5. Algorithm 5. + * chain(X, i, s, PK.seed, ADRS) + * 1: tmp <- X + * 2: for j from i to i + s - 1 do + * 3: ADRS.setHashAddress(j) + * 4: tmp <- F(PK.seed, ADRS, tmp + * 5: end for + * 6: return tmp + * + * @param [in, out] sk Eight hashes to iterate. + * @param [in] pk_seed Public key seed. + * @param [in] addr Encoded HashAddress. + * @param [in] ca Chain address start index. + * @param [in] n Number of bytes in each hash. + * @param [in] heap Dynamic memory allocation hint. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_chain_x8(byte* sk, const byte* pk_seed, byte* addr, + byte ca, byte n, void* heap) +{ + int ret = 0; + int j; + word32 o = 0; + /* Words the eight hashes occupy in the state. */ + word32 hw = (word32)(n / 8) * 8; + WC_DECLARE_VAR(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap); + + (void)heap; + + WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + if (ret == 0) { + o = slhdsakey_shake256_set_seed_ha_x8(fixed, pk_seed, addr, n); + slhdsakey_shake256_set_chain_addr_x8(fixed, o, ca); + slhdsakey_shake256_set_hash_x8(state, o, sk, n); + + for (j = 0; j < SLHDSA_WM1; j++) { + if (j != 0) { + /* Feed the previous output back in as the next input. */ + XMEMCPY(state + o, state, hw * sizeof(word64)); + } + XMEMCPY(state, fixed, o * sizeof(word64)); + slhdsakey_shake256_set_hash_addr_x8(state, o, (byte)j); + slhdsakey_shake256_set_end_x8(state, o + hw); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret != 0) + break; + sha3_blocksx8_avx512(state); + RESTORE_VECTOR_REGISTERS(); + } + + if (ret == 0) { + slhdsakey_shake256_get_hash_x8(state, sk, n); + } + } + + /* state holds the secret WOTS+ chain value; guard against a NULL state + * after an allocation failure (WOLFSSL_SMALL_STACK). */ + if (WC_VAR_OK(state)) { + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + } + WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, heap, DYNAMIC_TYPE_SLHDSA); + return ret; +} + +/* Generate the WOTS+ public key chains eight addresses at a time. + * + * FIPS 205 Section 5.1. Algorithm 6. + * wots_pkGen(SK.seed, PK.seed, ADRS) + * + * @param [in] key SLH-DSA key. + * @param [in] sk_seed Private key seed. + * @param [in] pk_seed Public key seed. + * @param [in] addr Encoded WOTS HASH HashAddress. + * @param [in] sk_addr Encoded WOTS PRF HashAddress. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, + const byte* pk_seed, byte* addr, byte* sk_addr) +{ + int ret = 0; + int i = 0; + byte n = key->params->n; + byte len = key->params->len; + WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N, + key->heap); + + WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N, + key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + for (i = 0; i < len - 7; i += 8) { + ret = slhdsakey_hash_prf_x8(pk_seed, sk_seed, sk_addr, n, (byte)i, + sk + i * n, key->heap); + if (ret != 0) { + break; + } + ret = slhdsakey_chain_x8(sk + i * n, pk_seed, addr, (byte)i, n, + key->heap); + if (ret != 0) { + break; + } + } + } + if (ret == 0) { + /* Trailing group, which runs past len into the buffer's slack. */ + ret = slhdsakey_hash_prf_x8(pk_seed, sk_seed, sk_addr, n, (byte)i, + sk + i * n, key->heap); + if (ret == 0) { + ret = slhdsakey_chain_x8(sk + i * n, pk_seed, addr, (byte)i, n, + key->heap); + } + } + if (ret == 0) { + ret = HASH_T_UPDATE(key, sk, (word32)len * n); + } + + /* On error sk still holds secret WOTS+ leaves, and the x8 PRF fills past + * len to an 8-lane multiple, so wipe the whole buffer. */ + if ((ret != 0) && WC_VAR_OK(sk)) { + ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N); + } + WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); + return ret; +} +#endif /* SLHDSA_HAVE_SHAKE_X8 */ + /* Generate WOTS+ public key - 4 consecutive addresses at a time. * * FIPS 205 Section 5.1. Algorithm 6. @@ -3261,6 +3661,13 @@ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, HA_Encode(sk_adrs, sk_addr); HA_Encode(adrs, addr); +#ifdef SLHDSA_HAVE_SHAKE_X8 + if (USE_INTEL_AVX512(cpuid_flags)) { + return slhdsakey_wots_pkgen_chain_x8(key, sk_seed, pk_seed, addr, + sk_addr); + } +#endif + #if !defined(WOLFSSL_SLHDSA_PARAM_NO_128) if (n == WC_SLHDSA_N_128) { ret = slhdsakey_wots_pkgen_chain_x4_16(key, sk_seed, pk_seed, addr, From 929f4c6823b868460780ad2e5597de3b0af406e2 Mon Sep 17 00:00:00 2001 From: Tobias Frauenschlaeger Date: Tue, 18 Aug 2026 10:41:32 +0000 Subject: [PATCH 3/6] Extend SLH-DSA eight-way hashing to FORS and cut its allocations The eight-way AVX512 Keccak path covered the WOTS+ chains only. Extend it to the FORS subtrees, which is where the remaining batched hashing sits. A FORS level narrower than eight nodes still finishes on the four-way path, so the last level of each subtree is unchanged. Separately, the batched helpers each allocated a Keccak state, used it for one group of hashes and freed it again. A WOTS+ public key needs one state, not one per group of four or eight, and a FORS subtree likewise. Hoist the state and the unchanging head of it into the callers that own the whole computation and pass them down. A SHAKE-128s signature under WOLFSSL_SMALL_STACK went from 143755 allocations to 11982, and the same change takes the buffers out of the callees' stack frames. Signing, measured on a Xeon at 2.1 GHz against the state before this work, taking each parameter set from one benchmark run: SHAKE-128s 344 ms -> 190 ms SHA2-128s 575 ms -> 484 ms SHAKE-192s 535 ms -> 315 ms SHA2-192s 1076 ms -> 816 ms SHAKE-256s 461 ms -> 274 ms SHA2-256s 907 ms -> 736 ms SHAKE-128f 15.5 ms -> 10.8 ms SHA2-128f 28.2 ms -> 22.0 ms SHAKE-192f 24.4 ms -> 17.1 ms SHA2-192f 45.8 ms -> 37.6 ms SHAKE-256f 51.5 ms -> 30.7 ms SHA2-256f 95.4 ms -> 76.5 ms SHA2 verification gains about a quarter from the hash rework. SHAKE verification is unchanged because it does not use the batched helpers. Public keys and signatures are unchanged. Checked byte for byte against the previous code for all twelve parameter sets, on the eight-way path and with NO_AVX512_SUPPORT forcing the four-way one. --- ChangeLog.md | 5 + wolfcrypt/src/wc_slhdsa.c | 809 +++++++++++++++++++++++--------------- 2 files changed, 491 insertions(+), 323 deletions(-) diff --git a/ChangeLog.md b/ChangeLog.md index fc1599cd74c..09d76ad9bdf 100644 --- a/ChangeLog.md +++ b/ChangeLog.md @@ -627,6 +627,11 @@ PR stands for Pull Request, and PR references a GitHub pull request num streaming two updates into it for every hash. by @Frauschi * Use the 8-way AVX512 Keccak permutation for the SLH-DSA WOTS+ chains on capable CPUs, which the four-way AVX2 path had to itself. by @Frauschi +* Extended the 8-way AVX512 Keccak permutation to the SLH-DSA FORS subtrees. + by @Frauschi +* Give each SLH-DSA WOTS+ public key and FORS subtree one Keccak state to + reuse, instead of allocating one per group of hashes. A SHAKE-128s + signature now makes about twelve times fewer allocations. by @Frauschi ## TLS/DTLS diff --git a/wolfcrypt/src/wc_slhdsa.c b/wolfcrypt/src/wc_slhdsa.c index b6dc3bcd444..1874176d38b 100644 --- a/wolfcrypt/src/wc_slhdsa.c +++ b/wolfcrypt/src/wc_slhdsa.c @@ -189,6 +189,20 @@ wc_static_assert(SLHDSA_MAX_MSG_SZ <= 255); #define SLHDSA_HAVE_SHAKE_X8 #endif +/* True when the eight-way path may be used for this CPU. */ +#ifdef SLHDSA_HAVE_SHAKE_X8 + #define SLHDSA_USE_SHAKE_X8() USE_INTEL_AVX512(cpuid_flags) +#else + #define SLHDSA_USE_SHAKE_X8() 0 +#endif + +/* Width of a Keccak state buffer that either path can use. */ +#ifdef SLHDSA_HAVE_SHAKE_X8 + #define SLHDSA_SHAKE_STATE_W SLHDSA_SHAKE_X8_STATE_W +#else + #define SLHDSA_SHAKE_STATE_W SLHDSA_SHAKE_X4_STATE_W +#endif + #ifndef WC_SLHDSA_ALL_NO_256F /* Maximum number of bytes to produce from digest of message. */ #define SLHDSA_MAX_MD 49 @@ -2538,32 +2552,24 @@ static int slhdsakey_chain_idx_x4_32(byte* sk, word32 i, word32 s, * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_hash_prf_x4(const byte* pk_seed, const byte* sk_seed, - byte* addr, byte n, byte ca, byte* sk, void* heap) + byte* addr, byte n, byte ca, byte* sk, word64* state) { - int ret = 0; - word32 o = 0; - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); - - (void)heap; + int ret; + word32 o; - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + o = slhdsakey_shake256_set_seed_ha_hash_x4(state, pk_seed, addr, sk_seed, + n); + SHAKE256_SET_CHAIN_ADDRESS(state, o, ca); + ret = SAVE_VECTOR_REGISTERS2(); if (ret == 0) { - o = slhdsakey_shake256_set_seed_ha_hash_x4(state, pk_seed, addr, - sk_seed, n); - SHAKE256_SET_CHAIN_ADDRESS(state, o, ca); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret == 0) { - sha3_blocksx4_avx2(state); - slhdsakey_shake256_get_hash_x4(state, sk, n); - RESTORE_VECTOR_REGISTERS(); - } - - /* state holds the secret PRF output (WOTS+ key). */ - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + sha3_blocksx4_avx2(state); + slhdsakey_shake256_get_hash_x4(state, sk, n); + RESTORE_VECTOR_REGISTERS(); } + /* state holds the secret PRF output (WOTS+ key). */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); + return ret; } @@ -2588,51 +2594,35 @@ static int slhdsakey_hash_prf_x4(const byte* pk_seed, const byte* sk_seed, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x4_16(byte* sk, const byte* pk_seed, byte* addr, - byte ca, void* heap) + byte ca, word64* fixed, word64* state) { int ret = 0; int j; - WC_DECLARE_VAR(fixed, word64, 8 * 4, heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); - (void)heap; + SHAKE256_SET_SEED_HA_X4_16(fixed, pk_seed, addr); + SHAKE256_SET_CHAIN_ADDRESS(fixed, 24, ca); + SHAKE256_SET_HASH_X4_16(state, sk); - WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, heap, DYNAMIC_TYPE_SLHDSA, - ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + for (j = 0; j < 15; j++) { + if (j != 0) { + XMEMCPY(state + 24, state, 16 * 4); + } + XMEMCPY(state, fixed, 24 * sizeof(word64)); + SHAKE256_SET_HASH_ADDRESS(state, 24, j); + SHAKE256_SET_END_X4(state, 32); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret != 0) + break; + sha3_blocksx4_avx2(state); + RESTORE_VECTOR_REGISTERS(); } - if (ret == 0) { - SHAKE256_SET_SEED_HA_X4_16(fixed, pk_seed, addr); - SHAKE256_SET_CHAIN_ADDRESS(fixed, 24, ca); - SHAKE256_SET_HASH_X4_16(state, sk); - for (j = 0; j < 15; j++) { - if (j != 0) { - XMEMCPY(state + 24, state, 16 * 4); - } - XMEMCPY(state, fixed, 24 * sizeof(word64)); - SHAKE256_SET_HASH_ADDRESS(state, 24, j); - SHAKE256_SET_END_X4(state, 32); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret != 0) - break; - sha3_blocksx4_avx2(state); - RESTORE_VECTOR_REGISTERS(); - } + if (ret == 0) + SHAKE256_GET_HASH_X4_16(state, sk); - if (ret == 0) - SHAKE256_GET_HASH_X4_16(state, sk); - } + /* state holds the secret WOTS+ chain value. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - /* state holds the secret WOTS+ chain value; guard against a NULL state - * after an allocation failure (WOLFSSL_SMALL_STACK). */ - if (WC_VAR_OK(state)) { - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - } - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif /* !WOLFSSL_SLHDSA_VERIFY_ONLY */ @@ -2658,51 +2648,35 @@ static int slhdsakey_chain_x4_16(byte* sk, const byte* pk_seed, byte* addr, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x4_24(byte* sk, const byte* pk_seed, byte* addr, - byte ca, void* heap) + byte ca, word64* fixed, word64* state) { int ret = 0; int j; - WC_DECLARE_VAR(fixed, word64, 8 * 4, heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); - (void)heap; + SHAKE256_SET_SEED_HA_X4_24(fixed, pk_seed, addr); + SHAKE256_SET_CHAIN_ADDRESS(fixed, 28, ca); + SHAKE256_SET_HASH_X4_24(state, sk); - WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, heap, DYNAMIC_TYPE_SLHDSA, - ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + for (j = 0; j < 15; j++) { + if (j != 0) { + XMEMCPY(state + 28, state, 24 * 4); + } + XMEMCPY(state, fixed, 28 * sizeof(word64)); + SHAKE256_SET_HASH_ADDRESS(state, 28, j); + SHAKE256_SET_END_X4(state, 40); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret != 0) + break; + sha3_blocksx4_avx2(state); + RESTORE_VECTOR_REGISTERS(); } - if (ret == 0) { - SHAKE256_SET_SEED_HA_X4_24(fixed, pk_seed, addr); - SHAKE256_SET_CHAIN_ADDRESS(fixed, 28, ca); - SHAKE256_SET_HASH_X4_24(state, sk); - for (j = 0; j < 15; j++) { - if (j != 0) { - XMEMCPY(state + 28, state, 24 * 4); - } - XMEMCPY(state, fixed, 28 * sizeof(word64)); - SHAKE256_SET_HASH_ADDRESS(state, 28, j); - SHAKE256_SET_END_X4(state, 40); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret != 0) - break; - sha3_blocksx4_avx2(state); - RESTORE_VECTOR_REGISTERS(); - } + if (ret == 0) + SHAKE256_GET_HASH_X4_24(state, sk); - if (ret == 0) - SHAKE256_GET_HASH_X4_24(state, sk); - } + /* state holds the secret WOTS+ chain value. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - /* state holds the secret WOTS+ chain value; guard against a NULL state - * after an allocation failure (WOLFSSL_SMALL_STACK). */ - if (WC_VAR_OK(state)) { - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - } - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif @@ -2728,51 +2702,35 @@ static int slhdsakey_chain_x4_24(byte* sk, const byte* pk_seed, byte* addr, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x4_32(byte* sk, const byte* pk_seed, byte* addr, - byte ca, void* heap) + byte ca, word64* fixed, word64* state) { int ret = 0; int j; - WC_DECLARE_VAR(fixed, word64, 8 * 4, heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); - (void)heap; + SHAKE256_SET_SEED_HA_X4_32(fixed, pk_seed, addr); + SHAKE256_SET_CHAIN_ADDRESS(fixed, 32, ca); + SHAKE256_SET_HASH_X4_32(state, sk); - WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, heap, DYNAMIC_TYPE_SLHDSA, - ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + for (j = 0; j < 15; j++) { + if (j != 0) { + XMEMCPY(state + 32, state, 32 * 4); + } + XMEMCPY(state, fixed, 32 * sizeof(word64)); + SHAKE256_SET_HASH_ADDRESS(state, 32, j); + SHAKE256_SET_END_X4(state, 48); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret != 0) + break; + sha3_blocksx4_avx2(state); + RESTORE_VECTOR_REGISTERS(); } - if (ret == 0) { - SHAKE256_SET_SEED_HA_X4_32(fixed, pk_seed, addr); - SHAKE256_SET_CHAIN_ADDRESS(fixed, 32, ca); - SHAKE256_SET_HASH_X4_32(state, sk); - for (j = 0; j < 15; j++) { - if (j != 0) { - XMEMCPY(state + 32, state, 32 * 4); - } - XMEMCPY(state, fixed, 32 * sizeof(word64)); - SHAKE256_SET_HASH_ADDRESS(state, 32, j); - SHAKE256_SET_END_X4(state, 48); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret != 0) - break; - sha3_blocksx4_avx2(state); - RESTORE_VECTOR_REGISTERS(); - } + if (ret == 0) + SHAKE256_GET_HASH_X4_32(state, sk); - if (ret == 0) - SHAKE256_GET_HASH_X4_32(state, sk); - } + /* state holds the secret WOTS+ chain value. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - /* state holds the secret WOTS+ chain value; guard against a NULL state - * after an allocation failure (WOLFSSL_SMALL_STACK). */ - if (WC_VAR_OK(state)) { - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - } - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif @@ -3110,18 +3068,30 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, int i; byte len = key->params->len; WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 16, key->heap); + WC_DECLARE_VAR(fixed, word64, 8 * 4, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap); + /* One state for every chain in this public key, rather than one per group + * of four. Each group refills it before use. */ WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 16, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, key->heap, DYNAMIC_TYPE_SLHDSA, + ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } if (ret == 0) { for (i = 0; i < len - 3; i += 4) { ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 16, (byte)i, - sk + i * 16, key->heap); + sk + i * 16, state); if (ret != 0) { break; } ret = slhdsakey_chain_x4_16(sk + i * 16, pk_seed, addr, (byte)i, - key->heap); + fixed, state); if (ret != 0) { break; } @@ -3129,10 +3099,10 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, } if (ret == 0) { ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 16, (byte)i, - sk + i * 16, key->heap); + sk + i * 16, state); if (ret == 0) { ret = slhdsakey_chain_x4_16(sk + i * 16, pk_seed, addr, (byte)i, - key->heap); + fixed, state); } } if (ret == 0) { @@ -3145,6 +3115,8 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, if ((ret != 0) && WC_VAR_OK(sk)) { ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 16); } + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } @@ -3184,18 +3156,30 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, int i; byte len = key->params->len; WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 24, key->heap); + WC_DECLARE_VAR(fixed, word64, 8 * 4, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap); + /* One state for every chain in this public key, rather than one per group + * of four. Each group refills it before use. */ WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 24, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, key->heap, DYNAMIC_TYPE_SLHDSA, + ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } if (ret == 0) { for (i = 0; i < len - 3; i += 4) { ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 24, (byte)i, - sk + i * 24, key->heap); + sk + i * 24, state); if (ret != 0) { break; } ret = slhdsakey_chain_x4_24(sk + i * 24, pk_seed, addr, (byte)i, - key->heap); + fixed, state); if (ret != 0) { break; } @@ -3203,10 +3187,10 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, } if (ret == 0) { ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 24, (byte)i, - sk + i * 24, key->heap); + sk + i * 24, state); if (ret == 0) { ret = slhdsakey_chain_x4_24(sk + i * 24, pk_seed, addr, (byte)i, - key->heap); + fixed, state); } } if (ret == 0) { @@ -3219,6 +3203,8 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, if ((ret != 0) && WC_VAR_OK(sk)) { ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 24); } + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } @@ -3258,18 +3244,30 @@ static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, int i; byte len = key->params->len; WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 32, key->heap); + WC_DECLARE_VAR(fixed, word64, 8 * 4, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap); + /* One state for every chain in this public key, rather than one per group + * of four. Each group refills it before use. */ WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 32, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, key->heap, DYNAMIC_TYPE_SLHDSA, + ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } if (ret == 0) { for (i = 0; i < len - 3; i += 4) { ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 32, (byte)i, - sk + i * 32, key->heap); + sk + i * 32, state); if (ret != 0) { break; } ret = slhdsakey_chain_x4_32(sk + i * 32, pk_seed, addr, (byte)i, - key->heap); + fixed, state); if (ret != 0) { break; } @@ -3277,10 +3275,10 @@ static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, } if (ret == 0) { ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 32, (byte)i, - sk + i * 32, key->heap); + sk + i * 32, state); if (ret == 0) { ret = slhdsakey_chain_x4_32(sk + i * 32, pk_seed, addr, (byte)i, - key->heap); + fixed, state); } } if (ret == 0) { @@ -3293,6 +3291,8 @@ static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, if ((ret != 0) && WC_VAR_OK(sk)) { ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 32); } + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } @@ -3432,6 +3432,55 @@ static void slhdsakey_shake256_set_hash_addr_x8(word64* state, word32 o, } } +/* Fill the 8-way state with the seed, encoded HashAddress and a hash that is + * the same in every lane. + * + * @param [out] state SHAKE-256 x8 state. + * @param [in] seed Seed at the start of each hash. + * @param [in] addr Encoded HashAddress for each hash. + * @param [in] hash Hash data to put into each hash. + * @param [in] n Number of bytes of seed. + * @return Offset after the seed and HashAddress, before the hash. + */ +static word32 slhdsakey_shake256_set_seed_ha_hash_x8(word64* state, + const byte* seed, const byte* addr, const byte* hash, int n) +{ + int i; + int l; + word32 o; + word32 ret; + + ret = o = slhdsakey_shake256_set_seed_ha_x8(state, seed, addr, n); + for (i = 0; i < n; i += 8) { + word64 v = readUnalignedWord64(hash + i); + + for (l = 0; l < 8; l++) { + state[o + l] = v; + } + o += 8; + } + slhdsakey_shake256_set_end_x8(state, o); + + return ret; +} + +/* Set an incrementing tree index into each lane of the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset of state after the HashAddress. + * @param [in] ti Value to encode, incrementing for each hash. + */ +static void slhdsakey_shake256_set_tree_index_x8(word64* state, word32 o, + word32 ti) +{ + int l; + + for (l = 0; l < 8; l++) { + c32toa(ti + (word32)l, + (byte*)&((word32*)(state + o - 8 + l))[1]); + } +} + /* PRF eight WOTS+ secret values at consecutive chain addresses. * * @param [in] pk_seed Public key seed. @@ -3440,50 +3489,29 @@ static void slhdsakey_shake256_set_hash_addr_x8(word64* state, word32 o, * @param [in] n Number of bytes in each hash. * @param [in] ca Chain address start index. * @param [out] sk Eight n-byte secret values. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned x8 Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_prf_x8(const byte* pk_seed, const byte* sk_seed, - byte* addr, byte n, byte ca, byte* sk, void* heap) + byte* addr, byte n, byte ca, byte* sk, word64* state) { - int ret = 0; + int ret; word32 o; - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap); - (void)heap; + o = slhdsakey_shake256_set_seed_ha_hash_x8(state, pk_seed, addr, sk_seed, + n); + slhdsakey_shake256_set_chain_addr_x8(state, o, ca); - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + ret = SAVE_VECTOR_REGISTERS2(); if (ret == 0) { - int i; - int l; - - o = slhdsakey_shake256_set_seed_ha_x8(state, pk_seed, addr, n); - slhdsakey_shake256_set_chain_addr_x8(state, o, ca); - /* PRF hashes the private key seed, the same value in every lane. */ - for (i = 0; i < n; i += 8) { - word64 v = readUnalignedWord64(sk_seed + i); - - for (l = 0; l < 8; l++) { - state[o + l] = v; - } - o += 8; - } - slhdsakey_shake256_set_end_x8(state, o); - - ret = SAVE_VECTOR_REGISTERS2(); - if (ret == 0) { - sha3_blocksx8_avx512(state); - slhdsakey_shake256_get_hash_x8(state, sk, n); - RESTORE_VECTOR_REGISTERS(); - } - - /* state holds the secret PRF output (WOTS+ key). */ - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + sha3_blocksx8_avx512(state); + slhdsakey_shake256_get_hash_x8(state, sk, n); + RESTORE_VECTOR_REGISTERS(); } + /* state holds the secret PRF output (WOTS+ key). */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + return ret; } @@ -3503,61 +3531,45 @@ static int slhdsakey_hash_prf_x8(const byte* pk_seed, const byte* sk_seed, * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. * @param [in] n Number of bytes in each hash. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] fixed Caller owned buffer for the unchanging state head. + * @param [in] state Caller owned x8 Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x8(byte* sk, const byte* pk_seed, byte* addr, - byte ca, byte n, void* heap) + byte ca, byte n, word64* fixed, word64* state) { int ret = 0; int j; - word32 o = 0; + word32 o; /* Words the eight hashes occupy in the state. */ word32 hw = (word32)(n / 8) * 8; - WC_DECLARE_VAR(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap); - (void)heap; - - WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - } - if (ret == 0) { - o = slhdsakey_shake256_set_seed_ha_x8(fixed, pk_seed, addr, n); - slhdsakey_shake256_set_chain_addr_x8(fixed, o, ca); - slhdsakey_shake256_set_hash_x8(state, o, sk, n); - - for (j = 0; j < SLHDSA_WM1; j++) { - if (j != 0) { - /* Feed the previous output back in as the next input. */ - XMEMCPY(state + o, state, hw * sizeof(word64)); - } - XMEMCPY(state, fixed, o * sizeof(word64)); - slhdsakey_shake256_set_hash_addr_x8(state, o, (byte)j); - slhdsakey_shake256_set_end_x8(state, o + hw); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret != 0) - break; - sha3_blocksx8_avx512(state); - RESTORE_VECTOR_REGISTERS(); - } + o = slhdsakey_shake256_set_seed_ha_x8(fixed, pk_seed, addr, n); + slhdsakey_shake256_set_chain_addr_x8(fixed, o, ca); + slhdsakey_shake256_set_hash_x8(state, o, sk, n); - if (ret == 0) { - slhdsakey_shake256_get_hash_x8(state, sk, n); + for (j = 0; j < SLHDSA_WM1; j++) { + if (j != 0) { + /* Feed the previous output back in as the next input. */ + XMEMCPY(state + o, state, hw * sizeof(word64)); } + XMEMCPY(state, fixed, o * sizeof(word64)); + slhdsakey_shake256_set_hash_addr_x8(state, o, (byte)j); + slhdsakey_shake256_set_end_x8(state, o + hw); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret != 0) + break; + sha3_blocksx8_avx512(state); + RESTORE_VECTOR_REGISTERS(); } - /* state holds the secret WOTS+ chain value; guard against a NULL state - * after an allocation failure (WOLFSSL_SMALL_STACK). */ - if (WC_VAR_OK(state)) { - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + if (ret == 0) { + slhdsakey_shake256_get_hash_x8(state, sk, n); } - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, heap, DYNAMIC_TYPE_SLHDSA); + + /* state holds the secret WOTS+ chain value. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + return ret; } @@ -3583,18 +3595,30 @@ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, byte len = key->params->len; WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N, key->heap); + WC_DECLARE_VAR(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, key->heap); + /* One state for every chain in this public key, rather than one per group + * of eight. Each group refills it before use. */ WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } if (ret == 0) { for (i = 0; i < len - 7; i += 8) { ret = slhdsakey_hash_prf_x8(pk_seed, sk_seed, sk_addr, n, (byte)i, - sk + i * n, key->heap); + sk + i * n, state); if (ret != 0) { break; } ret = slhdsakey_chain_x8(sk + i * n, pk_seed, addr, (byte)i, n, - key->heap); + fixed, state); if (ret != 0) { break; } @@ -3603,10 +3627,10 @@ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, if (ret == 0) { /* Trailing group, which runs past len into the buffer's slack. */ ret = slhdsakey_hash_prf_x8(pk_seed, sk_seed, sk_addr, n, (byte)i, - sk + i * n, key->heap); + sk + i * n, state); if (ret == 0) { ret = slhdsakey_chain_x8(sk + i * n, pk_seed, addr, (byte)i, n, - key->heap); + fixed, state); } } if (ret == 0) { @@ -3618,6 +3642,8 @@ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, if ((ret != 0) && WC_VAR_OK(sk)) { ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N); } + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } @@ -5556,32 +5582,24 @@ static int slhdsakey_fors_sk_gen(SlhDsaKey* key, const byte* sk_seed, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_prf_ti_x4(const byte* pk_seed, const byte* sk_seed, - byte* addr, byte n, word32 ti, byte* node, void* heap) + byte* addr, byte n, word32 ti, byte* node, word64* state) { - int ret = 0; - word32 o = 0; - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); - - (void)heap; + int ret; + word32 o; - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + o = slhdsakey_shake256_set_seed_ha_hash_x4(state, pk_seed, addr, sk_seed, + n); + SHAKE256_SET_TREE_INDEX(state, o, ti); + ret = SAVE_VECTOR_REGISTERS2(); if (ret == 0) { - o = slhdsakey_shake256_set_seed_ha_hash_x4(state, pk_seed, addr, - sk_seed, n); - SHAKE256_SET_TREE_INDEX(state, o, ti); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret == 0) { - sha3_blocksx4_avx2(state); - RESTORE_VECTOR_REGISTERS(); - slhdsakey_shake256_get_hash_x4(state, node, n); - } - - /* state holds the secret PRF output (FORS key). */ - ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + sha3_blocksx4_avx2(state); + RESTORE_VECTOR_REGISTERS(); + slhdsakey_shake256_get_hash_x4(state, node, n); } + /* state holds the secret PRF output (FORS key). */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); + return ret; } @@ -5606,36 +5624,27 @@ static int slhdsakey_hash_prf_ti_x4(const byte* pk_seed, const byte* sk_seed, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, - byte n, word32 ti, void* heap) + byte n, word32 ti, word64* state) { - int ret = 0; + int ret; int i; - word32 o = 0; - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); - - (void)heap; + word32 o; - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + o = slhdsakey_shake256_set_seed_ha_x4(state, pk_seed, addr, n); + SHAKE256_SET_TREE_INDEX(state, o, ti); + for (i = 0; i < n / 8; i++) { + state[o + 0] = readUnalignedWord64(node + 0 * n + i * 8); + state[o + 1] = readUnalignedWord64(node + 1 * n + i * 8); + state[o + 2] = readUnalignedWord64(node + 2 * n + i * 8); + state[o + 3] = readUnalignedWord64(node + 3 * n + i * 8); + o += 4; + } + SHAKE256_SET_END_X4(state, o); + ret = SAVE_VECTOR_REGISTERS2(); if (ret == 0) { - o = slhdsakey_shake256_set_seed_ha_x4(state, pk_seed, addr, n); - SHAKE256_SET_TREE_INDEX(state, o, ti); - for (i = 0; i < n / 8; i++) { - state[o + 0] = readUnalignedWord64(node + 0 * n + i * 8); - state[o + 1] = readUnalignedWord64(node + 1 * n + i * 8); - state[o + 2] = readUnalignedWord64(node + 2 * n + i * 8); - state[o + 3] = readUnalignedWord64(node + 3 * n + i * 8); - o += 4; - } - SHAKE256_SET_END_X4(state, o); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret == 0) { - sha3_blocksx4_avx2(state); - RESTORE_VECTOR_REGISTERS(); - slhdsakey_shake256_get_hash_x4(state, node, n); - } - - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + sha3_blocksx4_avx2(state); + RESTORE_VECTOR_REGISTERS(); + slhdsakey_shake256_get_hash_x4(state, node, n); } return ret; @@ -5663,40 +5672,144 @@ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_h_ti_x4(const byte* pk_seed, byte* addr, - const byte* m, byte n, word32 ti, byte* hash, void* heap) + const byte* m, byte n, word32 ti, byte* hash, word64* state) { - int ret = 0; + int ret; int i; - word32 o = 0; - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap); + word32 o; - (void)heap; + o = slhdsakey_shake256_set_seed_ha_x4(state, pk_seed, addr, n); + SHAKE256_SET_TREE_INDEX(state, o, ti); + for (i = 0; i < 2 * n / 8; i++) { + state[o + 0] = readUnalignedWord64(m + 0 * n + i * 8); + state[o + 1] = readUnalignedWord64(m + 2 * n + i * 8); + state[o + 2] = readUnalignedWord64(m + 4 * n + i * 8); + state[o + 3] = readUnalignedWord64(m + 6 * n + i * 8); + o += 4; + } + SHAKE256_SET_END_X4(state, o); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + sha3_blocksx4_avx2(state); + RESTORE_VECTOR_REGISTERS(); + slhdsakey_shake256_get_hash_x4(state, hash, n); + } - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + return ret; +} + +#ifdef SLHDSA_HAVE_SHAKE_X8 +/* PRF eight FORS secret values at consecutive tree indices. + * + * @param [in] pk_seed Public key seed. + * @param [in] sk_seed Private key seed. + * @param [in] addr Encoded HashAddress. + * @param [in] n Number of bytes in each hash. + * @param [in] ti Tree index start value. + * @param [out] node Eight n-byte outputs. + * @param [in] heap Dynamic memory allocation hint. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_hash_prf_ti_x8(const byte* pk_seed, const byte* sk_seed, + byte* addr, byte n, word32 ti, byte* node, word64* state) +{ + int ret; + word32 o; + + o = slhdsakey_shake256_set_seed_ha_hash_x8(state, pk_seed, addr, sk_seed, + n); + slhdsakey_shake256_set_tree_index_x8(state, o, ti); + ret = SAVE_VECTOR_REGISTERS2(); if (ret == 0) { - o = slhdsakey_shake256_set_seed_ha_x4(state, pk_seed, addr, n); - SHAKE256_SET_TREE_INDEX(state, o, ti); - for (i = 0; i < 2 * n / 8; i++) { - state[o + 0] = readUnalignedWord64(m + 0 * n + i * 8); - state[o + 1] = readUnalignedWord64(m + 2 * n + i * 8); - state[o + 2] = readUnalignedWord64(m + 4 * n + i * 8); - state[o + 3] = readUnalignedWord64(m + 6 * n + i * 8); - o += 4; - } - SHAKE256_SET_END_X4(state, o); - ret = SAVE_VECTOR_REGISTERS2(); - if (ret == 0) { - sha3_blocksx4_avx2(state); - RESTORE_VECTOR_REGISTERS(); - slhdsakey_shake256_get_hash_x4(state, hash, n); + sha3_blocksx8_avx512(state); + RESTORE_VECTOR_REGISTERS(); + slhdsakey_shake256_get_hash_x8(state, node, n); + } + + /* state holds the secret PRF output (FORS key). */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + + return ret; +} + +/* F hash eight at a time, varying by tree index. + * + * @param [in] pk_seed Public key seed. + * @param [in] addr Encoded HashAddress. + * @param [in, out] node On in, eight n-byte messages. On out, the outputs. + * @param [in] n Number of bytes in hash output. + * @param [in] ti Tree index start value. + * @param [in] heap Dynamic memory allocation hint. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_hash_f_ti_x8(const byte* pk_seed, byte* addr, byte* node, + byte n, word32 ti, word64* state) +{ + int ret; + int i; + int l; + word32 o; + + o = slhdsakey_shake256_set_seed_ha_x8(state, pk_seed, addr, n); + slhdsakey_shake256_set_tree_index_x8(state, o, ti); + for (i = 0; i < n / 8; i++) { + for (l = 0; l < 8; l++) { + state[o + l] = readUnalignedWord64(node + l * n + i * 8); } + o += 8; + } + slhdsakey_shake256_set_end_x8(state, o); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + sha3_blocksx8_avx512(state); + RESTORE_VECTOR_REGISTERS(); + slhdsakey_shake256_get_hash_x8(state, node, n); + } - WC_FREE_VAR_EX(state, heap, DYNAMIC_TYPE_SLHDSA); + return ret; +} + +/* H hash eight at a time, varying by tree index. + * + * @param [in] pk_seed Public key seed. + * @param [in] addr Encoded HashAddress. + * @param [in] m Sixteen n-byte values, a 2n-byte message per lane. + * @param [in] n Number of bytes in hash output. + * @param [in] ti Tree index start value. + * @param [out] hash Buffer to hold eight n-byte hash outputs. + * @param [in] heap Dynamic memory allocation hint. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_hash_h_ti_x8(const byte* pk_seed, byte* addr, + const byte* m, byte n, word32 ti, byte* hash, word64* state) +{ + int ret; + int i; + int l; + word32 o; + + o = slhdsakey_shake256_set_seed_ha_x8(state, pk_seed, addr, n); + slhdsakey_shake256_set_tree_index_x8(state, o, ti); + for (i = 0; i < 2 * n / 8; i++) { + for (l = 0; l < 8; l++) { + state[o + l] = readUnalignedWord64(m + 2 * l * n + i * 8); + } + o += 8; + } + slhdsakey_shake256_set_end_x8(state, o); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret == 0) { + sha3_blocksx8_avx512(state); + RESTORE_VECTOR_REGISTERS(); + slhdsakey_shake256_get_hash_x8(state, hash, n); } return ret; } +#endif /* SLHDSA_HAVE_SHAKE_X8 */ /* A ranges from 6-14. */ #if SLHDSA_MAX_A < 9 @@ -5857,12 +5970,14 @@ static int slhdsakey_fors_node_x4_z1(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] adrs FORS tree HashAddress. * @param [out] node n-byte root node. + * @param [in] state Batched Keccak state. * @return 0 on success. * @return SHAKE-256 error return code on digest failure. * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_fors_node_x4_low(SlhDsaKey* key, const byte* sk_seed, - word32 i, word32 z, const byte* pk_seed, word32* adrs, byte* node) + word32 i, word32 z, const byte* pk_seed, word32* adrs, byte* node, + word64* state) { int ret = 0; byte n = key->params->n; @@ -5889,25 +6004,44 @@ static int slhdsakey_fors_node_x4_low(SlhDsaKey* key, const byte* sk_seed, HA_Encode(adrs, addr); /* Step 2: Generate private key values for leaf indices. */ - for (j = 0; j < m; j += 4) { - ret = slhdsakey_hash_prf_ti_x4(pk_seed, sk_seed, sk_addr, n, - m * i + j, nodes + j * n, key->heap); - if (ret != 0) { - break; + j = 0; +#ifdef SLHDSA_HAVE_SHAKE_X8 + if (SLHDSA_USE_SHAKE_X8()) { + for (; j + 7 < m; j += 8) { + ret = slhdsakey_hash_prf_ti_x8(pk_seed, sk_seed, sk_addr, n, + m * i + j, nodes + j * n, state); + if (ret != 0) { + break; + } } } +#endif + /* Levels narrower than eight nodes finish four at a time. */ + for (; (ret == 0) && (j < m); j += 4) { + ret = slhdsakey_hash_prf_ti_x4(pk_seed, sk_seed, sk_addr, n, + m * i + j, nodes + j * n, state); + } } if (ret == 0) { /* Step 3: Set tree height to zero. */ HA_SetTreeHeight((word32*)addr, 0); /* Step 4-5: Set tree indices and compute leaf node. */ - for (j = 0; j < m; j += 4) { - ret = slhdsakey_hash_f_ti_x4(pk_seed, addr, nodes + j * n, n, - m * i + j, key->heap); - if (ret != 0) { - break; + j = 0; +#ifdef SLHDSA_HAVE_SHAKE_X8 + if (SLHDSA_USE_SHAKE_X8()) { + for (; j + 7 < m; j += 8) { + ret = slhdsakey_hash_f_ti_x8(pk_seed, addr, nodes + j * n, n, + m * i + j, state); + if (ret != 0) { + break; + } } } +#endif + for (; (ret == 0) && (j < m); j += 4) { + ret = slhdsakey_hash_f_ti_x4(pk_seed, addr, nodes + j * n, n, + m * i + j, state); + } } if (ret == 0) { word32 k; @@ -5916,13 +6050,23 @@ static int slhdsakey_fors_node_x4_low(SlhDsaKey* key, const byte* sk_seed, /* Step 9: Set tree height. */ HA_SetTreeHeightBE(addr, k); /* Step 10-11: Set tree index and compute nodes. */ - for (j = 0; j < m; j += 4) { - ret = slhdsakey_hash_h_ti_x4(pk_seed, addr, nodes + 2 * j * n, - n, m * i + j, nodes + j * n, key->heap); - if (ret != 0) { - break; + j = 0; +#ifdef SLHDSA_HAVE_SHAKE_X8 + if (SLHDSA_USE_SHAKE_X8()) { + for (; j + 7 < m; j += 8) { + ret = slhdsakey_hash_h_ti_x8(pk_seed, addr, + nodes + 2 * j * n, n, m * i + j, nodes + j * n, + state); + if (ret != 0) { + break; + } } } +#endif + for (; (ret == 0) && (j < m); j += 4) { + ret = slhdsakey_hash_h_ti_x4(pk_seed, addr, nodes + 2 * j * n, + n, m * i + j, nodes + j * n, state); + } if (ret != 0) { break; } @@ -5986,12 +6130,14 @@ static int slhdsakey_fors_node_x4_low(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] adrs FORS tree HashAddress. * @param [out] node n-byte root node. + * @param [in] state Batched Keccak state. * @return 0 on success. * @return SHAKE-256 error return code on digest failure. * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_fors_node_x4_high(SlhDsaKey* key, const byte* sk_seed, - word32 i, word32 z, const byte* pk_seed, word32* adrs, byte* node) + word32 i, word32 z, const byte* pk_seed, word32* adrs, byte* node, + word64* state) { int ret = 0; byte n = key->params->n; @@ -6011,7 +6157,7 @@ static int slhdsakey_fors_node_x4_high(SlhDsaKey* key, const byte* sk_seed, /* Steps 7-8: Compute left and right nodes. */ for (j = 0; j < m; j++) { ret = slhdsakey_fors_node_x4_low(key, sk_seed, m * i + j, z - z2, - pk_seed, adrs, nodes + j * n); + pk_seed, adrs, nodes + j * n, state); if (ret != 0) { break; } @@ -6028,13 +6174,23 @@ static int slhdsakey_fors_node_x4_high(SlhDsaKey* key, const byte* sk_seed, /* Encode FORS tree address for hashing. */ HA_Encode(adrs, addr); /* Step 10-11: Set tree index and compute nodes. */ - for (j = 0; j < m; j += 4) { - ret = slhdsakey_hash_h_ti_x4(pk_seed, addr, nodes + 2 * j * n, - n, m * i + j, nodes + j * n, key->heap); - if (ret != 0) { - break; + j = 0; +#ifdef SLHDSA_HAVE_SHAKE_X8 + if (SLHDSA_USE_SHAKE_X8()) { + for (; j + 7 < m; j += 8) { + ret = slhdsakey_hash_h_ti_x8(pk_seed, addr, + nodes + 2 * j * n, n, m * i + j, nodes + j * n, + state); + if (ret != 0) { + break; + } } } +#endif + for (; (ret == 0) && (j < m); j += 4) { + ret = slhdsakey_hash_h_ti_x4(pk_seed, addr, nodes + 2 * j * n, + n, m * i + j, nodes + j * n, state); + } if (ret != 0) { break; } @@ -6106,6 +6262,7 @@ static int slhdsakey_fors_node_x4(SlhDsaKey* key, const byte* sk_seed, word32 i, word32 z, const byte* pk_seed, word32* adrs, byte* node) { int ret = 0; + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_STATE_W, key->heap); /* Step 1: Check if we are at leaf node. */ if (z == 0) { @@ -6115,18 +6272,24 @@ static int slhdsakey_fors_node_x4(SlhDsaKey* key, const byte* sk_seed, word32 i, else if (z == 1) { ret = slhdsakey_fors_node_x4_z1(key, sk_seed, i, pk_seed, adrs, node); } - /* Step 6: 2-MAX_DEPTH levels above leaf node. */ - else if ((z >= 2) && (z <= SLHDSA_MAX_FORS_NODE_DEPTH)) { - ret = slhdsakey_fors_node_x4_low(key, sk_seed, i, z, pk_seed, adrs, - node); - } -#if SLHDSA_MAX_FORS_NODE_DEPTH < SLHDSA_MAX_A-1 - /* Step 6: More than MAX_DEPTH levels above leaf node. */ else { - ret = slhdsakey_fors_node_x4_high(key, sk_seed, i, z, pk_seed, adrs, - node); + /* One state for every batch of hashes in the subtree. */ + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + /* Step 6: 2-MAX_DEPTH levels above leaf node. */ + if ((ret == 0) && (z <= SLHDSA_MAX_FORS_NODE_DEPTH)) { + ret = slhdsakey_fors_node_x4_low(key, sk_seed, i, z, pk_seed, + adrs, node, state); + } + #if SLHDSA_MAX_FORS_NODE_DEPTH < SLHDSA_MAX_A-1 + /* Step 6: More than MAX_DEPTH levels above leaf node. */ + else if (ret == 0) { + ret = slhdsakey_fors_node_x4_high(key, sk_seed, i, z, pk_seed, + adrs, node, state); + } + #endif + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); } -#endif return ret; } From 481f6b873451635ce26843dcc58d34f1a8191df1 Mon Sep 17 00:00:00 2001 From: Tobias Frauenschlaeger Date: Tue, 18 Aug 2026 11:27:50 +0000 Subject: [PATCH 4/6] Batch SLH-DSA verification and reuse WOTS+ buffers per subtree Completing the WOTS+ chains of a signature was the last batched hashing still on the four-way AVX2 path, so verification saw none of the earlier work. Chains are gathered in order of the message value they start at, which means a group of eight has to be brought up in steps: each chain joins the ones already running when its own start is reached, and the whole group then grows to w-1 together. Written as a loop over the group rather than the unrolled ladder the four-way code uses. The last few chains, which do not fill a group, finish one at a time. The buffers a WOTS+ public key needs were still allocated per public key, and a subtree builds thousands of them. Hand them to the subtree instead: xmss_sign and key generation own one set and pass it down. A SHAKE-128s signature under WOLFSSL_SMALL_STACK now makes 1081 allocations where it made 143755 before this series, and 225 on the build without assembly. Peak stack is unchanged, since the buffers were live across the same calls either way. Measured on a Xeon at 2.1 GHz against the state before this series, each parameter set from one benchmark run: sign verify SHAKE-128s 344 -> 190 ms 0.370 -> 0.269 ms SHAKE-192s 535 -> 315 ms 0.523 -> 0.366 ms SHAKE-256s 461 -> 276 ms 0.743 -> 0.519 ms SHAKE-128f 15.5 -> 10.1 ms 1.039 -> 0.702 ms SHAKE-192f 24.4 -> 15.6 ms 1.433 -> 0.959 ms SHAKE-256f 51.5 -> 30.4 ms 1.467 -> 0.969 ms SHA2-128s 575 -> 469 ms 0.573 -> 0.479 ms SHA2-256f 95.4 -> 75.8 ms 2.456 -> 1.995 ms FORS verification and the WOTS+ signing chains stay on the four-way path. They are about a tenth and a five hundredth of their operations respectively, so batching them further is not worth the code. Public keys and signatures are unchanged, checked byte for byte against the previous code for all twelve parameter sets on both paths. Verification is the one path here where a defect means a false accept rather than a visible failure, so the check for that goes in the test suite rather than beside it. test_wc_slhdsa_sign_vfy now hands each signature it makes to a helper that requires SIG_VERIFY_E from nine single bit flips spread evenly over the signature, a changed message, a changed context and an absent context, then re-checks the restored signature so a rejection cannot be an artefact of leftover damage. Reusing the keys and signatures sign_vfy already makes keeps the cost to a few verifies per parameter set. Confirmed to bite by reverting the root comparison to always accept. --- ChangeLog.md | 5 + tests/api/test_slhdsa.c | 65 +++ wolfcrypt/src/wc_slhdsa.c | 885 ++++++++++++++++++++++++-------------- 3 files changed, 637 insertions(+), 318 deletions(-) diff --git a/ChangeLog.md b/ChangeLog.md index 09d76ad9bdf..ab31caa70f7 100644 --- a/ChangeLog.md +++ b/ChangeLog.md @@ -632,6 +632,11 @@ PR stands for Pull Request, and PR references a GitHub pull request num * Give each SLH-DSA WOTS+ public key and FORS subtree one Keccak state to reuse, instead of allocating one per group of hashes. A SHAKE-128s signature now makes about twelve times fewer allocations. by @Frauschi +* Use the 8-way AVX512 Keccak permutation when completing the WOTS+ chains of + an SLH-DSA signature, which speeds up verification. by @Frauschi +* Give a whole SLH-DSA XMSS subtree one set of WOTS+ buffers to reuse. A + SHAKE-128s signature now makes about 1000 allocations where it made about + 144000. by @Frauschi ## TLS/DTLS diff --git a/tests/api/test_slhdsa.c b/tests/api/test_slhdsa.c index dbe8c105683..da3e22acd2d 100644 --- a/tests/api/test_slhdsa.c +++ b/tests/api/test_slhdsa.c @@ -855,6 +855,47 @@ int test_wc_slhdsa_verify(void) /* * Test combined sign and verify for all parameter sets. */ +#if defined(WOLFSSL_HAVE_SLHDSA) && !defined(WOLFSSL_SLHDSA_VERIFY_ONLY) +/* Number of single bit flips applied across one signature. */ +#define TEST_SLHDSA_NEG_FLIPS 9 + +/* Verify has to reject a flipped signature bit, a changed message or context, + * and an absent context. The inputs are restored before returning. */ +static int slhdsa_verify_reject(SlhDsaKey* key, byte* ctx, byte ctxSz, + byte* msg, word32 msgSz, byte* sig, word32 sigLen) +{ + EXPECT_DECLS; + word32 off; + int i; + + for (i = 0; (i < TEST_SLHDSA_NEG_FLIPS) && EXPECT_SUCCESS(); i++) { + /* Evenly spaced, first and last byte included. */ + off = (word32)i * (sigLen - 1) / (TEST_SLHDSA_NEG_FLIPS - 1); + sig[off] ^= 0x01; + ExpectIntEQ(wc_SlhDsaKey_Verify(key, ctx, ctxSz, msg, msgSz, sig, + sigLen), WC_NO_ERR_TRACE(SIG_VERIFY_E)); + sig[off] ^= 0x01; + } + + msg[0] ^= 0x01; + ExpectIntEQ(wc_SlhDsaKey_Verify(key, ctx, ctxSz, msg, msgSz, sig, sigLen), + WC_NO_ERR_TRACE(SIG_VERIFY_E)); + msg[0] ^= 0x01; + + ctx[0] ^= 0x01; + ExpectIntEQ(wc_SlhDsaKey_Verify(key, ctx, ctxSz, msg, msgSz, sig, sigLen), + WC_NO_ERR_TRACE(SIG_VERIFY_E)); + ctx[0] ^= 0x01; + + ExpectIntEQ(wc_SlhDsaKey_Verify(key, NULL, 0, msg, msgSz, sig, sigLen), + WC_NO_ERR_TRACE(SIG_VERIFY_E)); + ExpectIntEQ(wc_SlhDsaKey_Verify(key, ctx, ctxSz, msg, msgSz, sig, sigLen), + 0); + + return EXPECT_RESULT(); +} +#endif /* WOLFSSL_HAVE_SLHDSA && !WOLFSSL_SLHDSA_VERIFY_ONLY */ + int test_wc_slhdsa_sign_vfy(void) { EXPECT_DECLS; @@ -886,6 +927,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHAKE128S_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif @@ -901,6 +944,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHAKE128F_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif @@ -916,6 +961,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, (word32)wc_SlhDsaKey_SigSize(&key)); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif @@ -931,6 +978,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHAKE192F_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif @@ -946,6 +995,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHAKE256S_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif @@ -961,6 +1012,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHAKE256F_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif @@ -976,6 +1029,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHA2_128S_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif #ifdef WOLFSSL_SLHDSA_PARAM_SHA2_128F @@ -988,6 +1043,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHA2_128F_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif #ifdef WOLFSSL_SLHDSA_PARAM_SHA2_192S @@ -1000,6 +1057,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHA2_192S_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif #ifdef WOLFSSL_SLHDSA_PARAM_SHA2_192F @@ -1012,6 +1071,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHA2_192F_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif #ifdef WOLFSSL_SLHDSA_PARAM_SHA2_256S @@ -1024,6 +1085,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHA2_256S_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif #ifdef WOLFSSL_SLHDSA_PARAM_SHA2_256F @@ -1036,6 +1099,8 @@ int test_wc_slhdsa_sign_vfy(void) ExpectIntEQ(sigLen, WC_SLHDSA_SHA2_256F_SIG_LEN); ExpectIntEQ(wc_SlhDsaKey_Verify(&key, ctx, sizeof(ctx), msg, sizeof(msg), sig, sigLen), 0); + ExpectIntEQ(slhdsa_verify_reject(&key, ctx, (byte)sizeof(ctx), msg, + (word32)sizeof(msg), sig, sigLen), TEST_SUCCESS); wc_SlhDsaKey_Free(&key); #endif #endif /* WOLFSSL_SLHDSA_SHA2 */ diff --git a/wolfcrypt/src/wc_slhdsa.c b/wolfcrypt/src/wc_slhdsa.c index 1874176d38b..fcae7ff31a1 100644 --- a/wolfcrypt/src/wc_slhdsa.c +++ b/wolfcrypt/src/wc_slhdsa.c @@ -203,6 +203,36 @@ wc_static_assert(SLHDSA_MAX_MSG_SZ <= 255); #define SLHDSA_SHAKE_STATE_W SLHDSA_SHAKE_X4_STATE_W #endif +/* Bytes of chain values one WOTS+ public key needs, including the overshoot + * the batched paths write past len. */ +#define SLHDSA_WOTS_SK_SZ ((SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N) + +/* A batched WOTS+ helper lets its last group of eight run past len, so the + * buffer has to cover the whole group holding index len - 1. */ +wc_static_assert(SLHDSA_WOTS_SK_SZ >= + (((SLHDSA_MAX_MSG_SZ + 7) / 8) * 8) * SLHDSA_MAX_N); + +/* Words in the unchanging head of the batched state, widest path in build. */ +#ifdef SLHDSA_HAVE_SHAKE_X8 + #define SLHDSA_WOTS_FIXED_W SLHDSA_SHAKE_X8_FIXED_W +#else + #define SLHDSA_WOTS_FIXED_W (8 * 4) +#endif + +/* Buffers reused by every WOTS+ public key computed from one subtree. A + * subtree builds thousands of them, so allocating per public key put the + * allocator on the hot path. */ +typedef struct SlhDsaWotsBufs { + /* Chain values of the public key being built. */ + byte* sk; +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + /* Unchanging head of the batched Keccak state. */ + word64* fixed; + /* Batched Keccak state. */ + word64* state; +#endif +} SlhDsaWotsBufs; + #ifndef WC_SLHDSA_ALL_NO_256F /* Maximum number of bytes to produce from digest of message. */ #define SLHDSA_MAX_MD 49 @@ -2527,6 +2557,183 @@ static int slhdsakey_chain_idx_x4_32(byte* sk, word32 i, word32 s, #endif #endif +#ifdef SLHDSA_HAVE_SHAKE_X8 +/* Fill the 8-way state with the seed and encoded HashAddress, one copy per + * lane. + * + * @param [out] state SHAKE-256 x8 state. + * @param [in] seed Seed at the start of each hash. + * @param [in] addr Encoded HashAddress for each hash. + * @param [in] n Number of bytes of seed. + * @return Offset after the seed and HashAddress. + */ +static word32 slhdsakey_shake256_set_seed_ha_x8(word64* state, + const byte* seed, const byte* addr, int n) +{ + int i; + int l; + word32 o = 0; + + for (i = 0; i < n; i += 8) { + word64 v = readUnalignedWord64(seed + i); + + for (l = 0; l < 8; l++) { + state[o + l] = v; + } + o += 8; + } + for (i = 0; i < SLHDSA_HA_SZ; i += 8) { + word64 v = readUnalignedWord64(addr + i); + + for (l = 0; l < 8; l++) { + state[o + l] = v; + } + o += 8; + } + + return o; +} + +/* Append one n-byte hash per lane to the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset to place the hashes at. + * @param [in] hash Eight n-byte hashes. + * @param [in] n Number of bytes in each hash. + */ +static void slhdsakey_shake256_set_hash_x8(word64* state, word32 o, + const byte* hash, int n) +{ + int i; + int l; + + for (i = 0; i < n; i += 8) { + for (l = 0; l < 8; l++) { + state[o + l] = readUnalignedWord64(hash + l * n + i); + } + o += 8; + } +} + +/* Get the eight SHAKE-256 n-byte hash results. + * + * @param [in] state SHAKE-256 x8 state. + * @param [out] hash Buffer to hold eight n-byte hash results. + * @param [in] n Length of each hash in bytes. + */ +static void slhdsakey_shake256_get_hash_x8(const word64* state, byte* hash, + int n) +{ + int i; + int l; + + for (i = 0; i < (n / 8); i++) { + for (l = 0; l < 8; l++) { + writeUnalignedWord64(hash + l * n + i * 8, state[8 * i + l]); + } + } +} + +/* Set the end of the SHAKE-256 x8 state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset to the end of the data. + */ +static void slhdsakey_shake256_set_end_x8(word64* state, word32 o) +{ + int l; + + /* Data end marker. */ + for (l = 0; l < 8; l++) { + state[o + l] = (word64)0x1f; + } + XMEMSET(state + o + 8, 0, + (size_t)(SLHDSA_SHAKE_X8_STATE_W - (o + 8)) * sizeof(word64)); + /* SHAKE-256 state end marker. */ + for (l = 0; l < 8; l++) { + ((word8*)(state + 8 * WC_SHA3_256_COUNT - 8 + l))[7] ^= 0x80; + } +} + +/* Set the same hash address into each lane of the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset of state after the HashAddress. + * @param [in] a Value to set for each hash. + */ +static void slhdsakey_shake256_set_hash_addr_x8(word64* state, word32 o, + byte a) +{ + int l; + + for (l = 0; l < 8; l++) { + ((word8*)(state + o - 8 + l))[7] = (word8)a; + } +} + +/* Set the chain address indices into each lane of the 8-way state. + * + * @param [in, out] state SHAKE-256 x8 state. + * @param [in] o Offset of state after the HashAddress. + * @param [in] idx Chain address to set for each hash. + */ +static void slhdsakey_shake256_set_chain_addr_idx_x8(word64* state, word32 o, + const byte* idx) +{ + int l; + + for (l = 0; l < 8; l++) { + ((word8*)(state + o - 8 + l))[3] = idx[l]; + } +} + +/* Iterate the hash function over eight chains with their own addresses. + * + * FIPS 205. Section 5. Algorithm 5. + * chain(X, i, s, PK.seed, ADRS) + * + * @param [in, out] sk Eight hashes to iterate. + * @param [in] i Step to start at. + * @param [in] s Number of steps to take. + * @param [in] n Number of bytes in each hash. + * @param [in] o Offset of the state after the seed and address. + * @param [in] fixed Caller owned state head, already filled. + * @param [in] state Caller owned x8 Keccak state. + * @return 0 on success. + */ +static int slhdsakey_chain_idx_x8(byte* sk, word32 i, word32 s, byte n, + word32 o, word64* fixed, word64* state) +{ + int ret = 0; + word32 j; + /* Words the eight hashes occupy in the state. */ + word32 hw = (word32)(n / 8) * 8; + + slhdsakey_shake256_set_hash_x8(state, o, sk, n); + + for (j = i; j < i + s; j++) { + if (j != i) { + XMEMCPY(state + o, state, hw * sizeof(word64)); + } + XMEMCPY(state, fixed, o * sizeof(word64)); + slhdsakey_shake256_set_hash_addr_x8(state, o, (byte)j); + slhdsakey_shake256_set_end_x8(state, o + hw); + ret = SAVE_VECTOR_REGISTERS2(); + if (ret != 0) + break; + sha3_blocksx8_avx512(state); + RESTORE_VECTOR_REGISTERS(); + } + + if (ret == 0) { + slhdsakey_shake256_get_hash_x8(state, sk, n); + } + + return ret; +} + +#endif /* SLHDSA_HAVE_SHAKE_X8 */ + #ifndef WOLFSSL_SLHDSA_VERIFY_ONLY #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) /* PRF hash 4 simultaneously. @@ -3062,39 +3269,26 @@ static int slhdsakey_chain_idx_32(SlhDsaKey* key, byte* sk, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, byte* addr, byte* sk_addr) + const byte* pk_seed, byte* addr, byte* sk_addr, + SlhDsaWotsBufs* bufs) { int ret = 0; int i; byte len = key->params->len; - WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 16, key->heap); - WC_DECLARE_VAR(fixed, word64, 8 * 4, key->heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap); + byte* sk = bufs->sk; + word64* fixed = bufs->fixed; + word64* state = bufs->state; - /* One state for every chain in this public key, rather than one per group - * of four. Each group refills it before use. */ - WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 16, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, key->heap, DYNAMIC_TYPE_SLHDSA, - ret = MEMORY_E); - } - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - } - if (ret == 0) { - for (i = 0; i < len - 3; i += 4) { - ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 16, (byte)i, - sk + i * 16, state); - if (ret != 0) { - break; - } - ret = slhdsakey_chain_x4_16(sk + i * 16, pk_seed, addr, (byte)i, - fixed, state); - if (ret != 0) { - break; - } + for (i = 0; i < len - 3; i += 4) { + ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 16, (byte)i, + sk + i * 16, state); + if (ret != 0) { + break; + } + ret = slhdsakey_chain_x4_16(sk + i * 16, pk_seed, addr, (byte)i, + fixed, state); + if (ret != 0) { + break; } } if (ret == 0) { @@ -3112,12 +3306,9 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, /* On error sk still holds secret WOTS+ leaves; on success it is overwritten * with public chain values. The x4 PRF fills up to a 4-lane multiple * (beyond len), so wipe the whole buffer. */ - if ((ret != 0) && WC_VAR_OK(sk)) { + if (ret != 0) { ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 16); } - WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif @@ -3150,39 +3341,26 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, byte* addr, byte* sk_addr) + const byte* pk_seed, byte* addr, byte* sk_addr, + SlhDsaWotsBufs* bufs) { int ret = 0; int i; byte len = key->params->len; - WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 24, key->heap); - WC_DECLARE_VAR(fixed, word64, 8 * 4, key->heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap); + byte* sk = bufs->sk; + word64* fixed = bufs->fixed; + word64* state = bufs->state; - /* One state for every chain in this public key, rather than one per group - * of four. Each group refills it before use. */ - WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 24, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, key->heap, DYNAMIC_TYPE_SLHDSA, - ret = MEMORY_E); - } - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - } - if (ret == 0) { - for (i = 0; i < len - 3; i += 4) { - ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 24, (byte)i, - sk + i * 24, state); - if (ret != 0) { - break; - } - ret = slhdsakey_chain_x4_24(sk + i * 24, pk_seed, addr, (byte)i, - fixed, state); - if (ret != 0) { - break; - } + for (i = 0; i < len - 3; i += 4) { + ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 24, (byte)i, + sk + i * 24, state); + if (ret != 0) { + break; + } + ret = slhdsakey_chain_x4_24(sk + i * 24, pk_seed, addr, (byte)i, + fixed, state); + if (ret != 0) { + break; } } if (ret == 0) { @@ -3200,12 +3378,9 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, /* On error sk still holds secret WOTS+ leaves; on success it is overwritten * with public chain values. The x4 PRF fills up to a 4-lane multiple * (beyond len), so wipe the whole buffer. */ - if ((ret != 0) && WC_VAR_OK(sk)) { + if (ret != 0) { ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 24); } - WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif @@ -3234,172 +3409,54 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] sk_addr Encoded WOTS PRF HashAddress. - * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. - */ -static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, byte* addr, byte* sk_addr) -{ - int ret = 0; - int i; - byte len = key->params->len; - WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 32, key->heap); - WC_DECLARE_VAR(fixed, word64, 8 * 4, key->heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap); - - /* One state for every chain in this public key, rather than one per group - * of four. Each group refills it before use. */ - WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * 32, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(fixed, word64, 8 * 4, key->heap, DYNAMIC_TYPE_SLHDSA, - ret = MEMORY_E); - } - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X4_STATE_W, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - } - if (ret == 0) { - for (i = 0; i < len - 3; i += 4) { - ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 32, (byte)i, - sk + i * 32, state); - if (ret != 0) { - break; - } - ret = slhdsakey_chain_x4_32(sk + i * 32, pk_seed, addr, (byte)i, - fixed, state); - if (ret != 0) { - break; - } - } - } - if (ret == 0) { - ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 32, (byte)i, - sk + i * 32, state); - if (ret == 0) { - ret = slhdsakey_chain_x4_32(sk + i * 32, pk_seed, addr, (byte)i, - fixed, state); - } - } - if (ret == 0) { - ret = HASH_T_UPDATE(key, sk, (word32)len * 32U); - } - - /* On error sk still holds secret WOTS+ leaves; on success it is overwritten - * with public chain values. The x4 PRF fills up to a 4-lane multiple - * (beyond len), so wipe the whole buffer. */ - if ((ret != 0) && WC_VAR_OK(sk)) { - ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 32); - } - WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); - return ret; -} -#endif - -#ifdef SLHDSA_HAVE_SHAKE_X8 -/* Fill the 8-way state with the seed and encoded HashAddress, one copy per - * lane. - * - * The 8-way helpers loop over the lanes rather than unrolling like the 4-way - * ones. Setting up the state is a fraction of a percent of a signature; the - * Keccak permutation is the rest. - * - * @param [out] state SHAKE-256 x8 state. - * @param [in] seed Seed at the start of each hash. - * @param [in] addr Encoded HashAddress for each hash. - * @param [in] n Number of bytes of seed. - * @return Offset after the seed and HashAddress. - */ -static word32 slhdsakey_shake256_set_seed_ha_x8(word64* state, - const byte* seed, const byte* addr, int n) -{ - int i; - int l; - word32 o = 0; - - for (i = 0; i < n; i += 8) { - word64 v = readUnalignedWord64(seed + i); - - for (l = 0; l < 8; l++) { - state[o + l] = v; - } - o += 8; - } - for (i = 0; i < SLHDSA_HA_SZ; i += 8) { - word64 v = readUnalignedWord64(addr + i); - - for (l = 0; l < 8; l++) { - state[o + l] = v; - } - o += 8; - } - - return o; -} - -/* Append one n-byte hash per lane to the 8-way state. - * - * @param [in, out] state SHAKE-256 x8 state. - * @param [in] o Offset to place the hashes at. - * @param [in] hash Eight n-byte hashes. - * @param [in] n Number of bytes in each hash. - */ -static void slhdsakey_shake256_set_hash_x8(word64* state, word32 o, - const byte* hash, int n) -{ - int i; - int l; - - for (i = 0; i < n; i += 8) { - for (l = 0; l < 8; l++) { - state[o + l] = readUnalignedWord64(hash + l * n + i); - } - o += 8; - } -} - -/* Get the eight SHAKE-256 n-byte hash results. - * - * @param [in] state SHAKE-256 x8 state. - * @param [out] hash Buffer to hold eight n-byte hash results. - * @param [in] n Length of each hash in bytes. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. */ -static void slhdsakey_shake256_get_hash_x8(const word64* state, byte* hash, - int n) +static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, + const byte* pk_seed, byte* addr, byte* sk_addr, + SlhDsaWotsBufs* bufs) { + int ret = 0; int i; - int l; + byte len = key->params->len; + byte* sk = bufs->sk; + word64* fixed = bufs->fixed; + word64* state = bufs->state; - for (i = 0; i < (n / 8); i++) { - for (l = 0; l < 8; l++) { - writeUnalignedWord64(hash + l * n + i * 8, state[8 * i + l]); + for (i = 0; i < len - 3; i += 4) { + ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 32, (byte)i, + sk + i * 32, state); + if (ret != 0) { + break; + } + ret = slhdsakey_chain_x4_32(sk + i * 32, pk_seed, addr, (byte)i, + fixed, state); + if (ret != 0) { + break; } } -} - -/* Set the end of the SHAKE-256 x8 state. - * - * @param [in, out] state SHAKE-256 x8 state. - * @param [in] o Offset to the end of the data. - */ -static void slhdsakey_shake256_set_end_x8(word64* state, word32 o) -{ - int l; - - /* Data end marker. */ - for (l = 0; l < 8; l++) { - state[o + l] = (word64)0x1f; + if (ret == 0) { + ret = slhdsakey_hash_prf_x4(pk_seed, sk_seed, sk_addr, 32, (byte)i, + sk + i * 32, state); + if (ret == 0) { + ret = slhdsakey_chain_x4_32(sk + i * 32, pk_seed, addr, (byte)i, + fixed, state); + } } - XMEMSET(state + o + 8, 0, - (size_t)(SLHDSA_SHAKE_X8_STATE_W - (o + 8)) * sizeof(word64)); - /* SHAKE-256 state end marker. */ - for (l = 0; l < 8; l++) { - ((word8*)(state + 8 * WC_SHA3_256_COUNT - 8 + l))[7] ^= 0x80; + if (ret == 0) { + ret = HASH_T_UPDATE(key, sk, (word32)len * 32U); + } + + /* On error sk still holds secret WOTS+ leaves, and the x4 PRF fills past + * len to a 4-lane multiple, so wipe the whole buffer. */ + if (ret != 0) { + ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 32); } + return ret; } +#endif +#ifdef SLHDSA_HAVE_SHAKE_X8 /* Set an incrementing chain address into each lane of the 8-way state. * * @param [in, out] state SHAKE-256 x8 state. @@ -3416,22 +3473,6 @@ static void slhdsakey_shake256_set_chain_addr_x8(word64* state, word32 o, } } -/* Set the same hash address into each lane of the 8-way state. - * - * @param [in, out] state SHAKE-256 x8 state. - * @param [in] o Offset of state after the HashAddress. - * @param [in] a Value to set for each hash. - */ -static void slhdsakey_shake256_set_hash_addr_x8(word64* state, word32 o, - byte a) -{ - int l; - - for (l = 0; l < 8; l++) { - ((word8*)(state + o - 8 + l))[7] = (word8)a; - } -} - /* Fill the 8-way state with the seed, encoded HashAddress and a hash that is * the same in every lane. * @@ -3587,41 +3628,27 @@ static int slhdsakey_chain_x8(byte* sk, const byte* pk_seed, byte* addr, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, byte* addr, byte* sk_addr) + const byte* pk_seed, byte* addr, byte* sk_addr, + SlhDsaWotsBufs* bufs) { int ret = 0; int i = 0; byte n = key->params->n; byte len = key->params->len; - WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N, - key->heap); - WC_DECLARE_VAR(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, key->heap); - WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, key->heap); + byte* sk = bufs->sk; + word64* fixed = bufs->fixed; + word64* state = bufs->state; - /* One state for every chain in this public key, rather than one per group - * of eight. Each group refills it before use. */ - WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N, - key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - if (ret == 0) { - WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - } - if (ret == 0) { - WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, key->heap, - DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - } - if (ret == 0) { - for (i = 0; i < len - 7; i += 8) { - ret = slhdsakey_hash_prf_x8(pk_seed, sk_seed, sk_addr, n, (byte)i, - sk + i * n, state); - if (ret != 0) { - break; - } - ret = slhdsakey_chain_x8(sk + i * n, pk_seed, addr, (byte)i, n, - fixed, state); - if (ret != 0) { - break; - } + for (i = 0; i < len - 7; i += 8) { + ret = slhdsakey_hash_prf_x8(pk_seed, sk_seed, sk_addr, n, (byte)i, + sk + i * n, state); + if (ret != 0) { + break; + } + ret = slhdsakey_chain_x8(sk + i * n, pk_seed, addr, (byte)i, n, + fixed, state); + if (ret != 0) { + break; } } if (ret == 0) { @@ -3639,12 +3666,9 @@ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, /* On error sk still holds secret WOTS+ leaves, and the x8 PRF fills past * len to an 8-lane multiple, so wipe the whole buffer. */ - if ((ret != 0) && WC_VAR_OK(sk)) { - ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 7) * SLHDSA_MAX_N); + if (ret != 0) { + ForceZero(sk, SLHDSA_WOTS_SK_SZ); } - WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); - WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif /* SLHDSA_HAVE_SHAKE_X8 */ @@ -3676,7 +3700,7 @@ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, word32* adrs, word32* sk_adrs) + const byte* pk_seed, word32* adrs, word32* sk_adrs, SlhDsaWotsBufs* bufs) { int ret = 0; byte sk_addr[SLHDSA_HA_SZ]; @@ -3690,28 +3714,28 @@ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, #ifdef SLHDSA_HAVE_SHAKE_X8 if (USE_INTEL_AVX512(cpuid_flags)) { return slhdsakey_wots_pkgen_chain_x8(key, sk_seed, pk_seed, addr, - sk_addr); + sk_addr, bufs); } #endif #if !defined(WOLFSSL_SLHDSA_PARAM_NO_128) if (n == WC_SLHDSA_N_128) { ret = slhdsakey_wots_pkgen_chain_x4_16(key, sk_seed, pk_seed, addr, - sk_addr); + sk_addr, bufs); } else #endif #if !defined(WOLFSSL_SLHDSA_PARAM_NO_192) if (n == 24) { ret = slhdsakey_wots_pkgen_chain_x4_24(key, sk_seed, pk_seed, addr, - sk_addr); + sk_addr, bufs); } else #endif #if !defined(WOLFSSL_SLHDSA_PARAM_NO_256) if (n == 32) { ret = slhdsakey_wots_pkgen_chain_x4_32(key, sk_seed, pk_seed, addr, - sk_addr); + sk_addr, bufs); } else #endif @@ -3751,7 +3775,8 @@ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, word32* adrs, word32* sk_adrs) + const byte* pk_seed, word32* adrs, word32* sk_adrs, + SlhDsaWotsBufs* bufs) { int ret = 0; int i; @@ -3759,12 +3784,9 @@ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, byte len = key->params->len; #if !defined(WOLFSSL_WC_SLHDSA_SMALL_MEM) - WC_DECLARE_VAR(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * SLHDSA_MAX_N, key->heap); + byte* sk = bufs->sk; - WC_ALLOC_VAR_EX(sk, byte, (SLHDSA_MAX_MSG_SZ + 3) * SLHDSA_MAX_N, - key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); - if (ret == 0) - XMEMSET(sk, 0, (SLHDSA_MAX_MSG_SZ + 3) * SLHDSA_MAX_N); + XMEMSET(sk, 0, SLHDSA_WOTS_SK_SZ); if (ret == 0) { /* Step 4. len consecutive addresses. */ for (i = 0; i < len; i++) { @@ -3792,13 +3814,14 @@ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, } /* On error sk still holds secret WOTS+ leaves; on success it is overwritten * with public chain values (generic path fills exactly len entries). */ - if ((ret != 0) && WC_VAR_OK(sk)) { + if (ret != 0) { ForceZero(sk, (word32)len * n); } - WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); #else byte sk[SLHDSA_MAX_N]; + (void)bufs; + /* Step 4. len consecutive addresses. */ for (i = 0; i < len; i++) { /* Step 5. Set chain address for WOTS PRF. */ @@ -3862,7 +3885,7 @@ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_wots_pkgen(SlhDsaKey* key, const byte* sk_seed, - const byte* pk_seed, word32* adrs, byte* node) + const byte* pk_seed, word32* adrs, byte* node, SlhDsaWotsBufs* bufs) { int ret; byte n = key->params->n; @@ -3891,14 +3914,14 @@ static int slhdsakey_wots_pkgen(SlhDsaKey* key, const byte* sk_seed, IS_INTEL_AVX2(cpuid_flags) && (SAVE_VECTOR_REGISTERS2() == 0)) { ret = slhdsakey_wots_pkgen_chain_x4(key, sk_seed, pk_seed, adrs, - sk_adrs); + sk_adrs, bufs); RESTORE_VECTOR_REGISTERS(); } else #endif { ret = slhdsakey_wots_pkgen_chain_c(key, sk_seed, pk_seed, adrs, - sk_adrs); + sk_adrs, bufs); } } if (ret == 0) { @@ -4628,6 +4651,123 @@ static int slhdsakey_chain_idx_to_max_32(SlhDsaKey* key, const byte* sig, * @return 0 on success. * @return MEMORY_E on dynamic memory allocation failure. */ +#ifdef SLHDSA_HAVE_SHAKE_X8 +/* Complete the WOTS+ chains of a signature, eight at a time. + * + * FIPS 205. Section 5.3. Algorithm 8. + * wots_pkFromSig(sig, M, PK.seed, ADRS) + * 10: tmp[i] <- chain(sig[i], msg[i], w - 1 - msg[i], PK.seed, ADRS) + * + * A group of eight is brought up in steps: each chain joins the ones already + * running when its own start value is reached, then the group grows to w-1 + * together. Lanes not yet carrying a signature value hash buffer contents and + * are overwritten before they matter. + * + * @param [in] key SLH-DSA key. + * @param [in] sig Signature - (2.n + 3) hashes of length n. + * @param [in] msg Encoded message with checksum. + * @param [in] pk_seed Public key seed. + * @param [in] adrs WOTS HASH HashAddress. + * @param [out] nodes Buffer to hold the len completed chains. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_wots_pk_from_sig_x8(SlhDsaKey* key, const byte* sig, + const byte* msg, const byte* pk_seed, word32* adrs, byte* nodes) +{ + int ret = 0; + int i; + int j; + int k; + word32 o = 0; + byte n = key->params->n; + byte len = key->params->len; + byte ii = 0; + byte idx[8]; + byte addr[SLHDSA_HA_SZ]; + WC_DECLARE_VAR(node, byte, 8 * SLHDSA_MAX_N, key->heap); + WC_DECLARE_VAR(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_X8_STATE_W, key->heap); + + XMEMSET(idx, 0, sizeof(idx)); + + WC_ALLOC_VAR_EX(node, byte, 8 * SLHDSA_MAX_N, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_SHAKE_X8_FIXED_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_X8_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + if (ret == 0) { + XMEMSET(node, 0, 8 * (size_t)n); + } + + for (j = 0; (ret == 0) && (j <= (int)SLHDSA_WM1); j++) { + for (i = 0; (ret == 0) && (i < (int)len); i++) { + if (msg[i] != (byte)j) { + continue; + } + idx[ii++] = (byte)i; + if (ii < 8) { + continue; + } + + HA_SetChainAddress(adrs, idx[0]); + HA_Encode(adrs, addr); + /* Seed, address and chain indices are the same for the group. */ + o = slhdsakey_shake256_set_seed_ha_x8(fixed, pk_seed, addr, n); + slhdsakey_shake256_set_chain_addr_idx_x8(fixed, o, idx); + /* Bring each chain up to the start of the next one. */ + for (k = 0; (ret == 0) && (k < 8); k++) { + XMEMCPY(node + k * n, sig + idx[k] * n, n); + if (k + 1 < 8) { + byte cur = msg[idx[k]]; + byte nxt = msg[idx[k + 1]]; + + if (cur != nxt) { + ret = slhdsakey_chain_idx_x8(node, cur, + (word32)(nxt - cur), n, o, fixed, state); + } + } + } + /* Grow the whole group to the end of the chain. */ + if ((ret == 0) && (j != (int)SLHDSA_WM1)) { + ret = slhdsakey_chain_idx_x8(node, (word32)j, + (word32)((int)SLHDSA_WM1 - j), n, o, fixed, state); + } + if (ret == 0) { + for (k = 0; k < 8; k++) { + XMEMCPY(nodes + idx[k] * n, node + k * n, n); + } + } + ii = 0; + } + } + + /* Chains left over after the last full group finish one at a time. */ + for (k = 0; (ret == 0) && (k < (int)ii); k++) { + HA_SetChainAddress(adrs, idx[k]); + XMEMCPY(node, sig + idx[k] * n, n); + ret = slhdsakey_chain(key, node, msg[idx[k]], + (byte)((int)SLHDSA_WM1 - msg[idx[k]]), pk_seed, adrs, node); + if (ret == 0) { + XMEMCPY(nodes + idx[k] * n, node, n); + } + } + + if (WC_VAR_OK(state)) { + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + } + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(node, key->heap, DYNAMIC_TYPE_SLHDSA); + return ret; +} +#endif /* SLHDSA_HAVE_SHAKE_X8 */ + static int slhdsakey_wots_pk_from_sig_x4(SlhDsaKey* key, const byte* sig, const byte* msg, const byte* pk_seed, word32* adrs, byte* pk_sig) { @@ -4640,6 +4780,13 @@ static int slhdsakey_wots_pk_from_sig_x4(SlhDsaKey* key, const byte* sig, WC_ALLOC_VAR_EX(nodes, byte, SLHDSA_MAX_MSG_SZ * SLHDSA_MAX_N, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); +#ifdef SLHDSA_HAVE_SHAKE_X8 + if ((ret == 0) && SLHDSA_USE_SHAKE_X8()) { + ret = slhdsakey_wots_pk_from_sig_x8(key, sig, msg, pk_seed, adrs, + nodes); + } + else +#endif #if !defined(WOLFSSL_SLHDSA_PARAM_NO_128) if ((ret == 0) && (n == WC_SLHDSA_N_128)) { int i; @@ -5017,7 +5164,8 @@ static int slhdsakey_wots_pk_from_sig(SlhDsaKey* key, const byte* sig, * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, - int z, const byte* pk_seed, word32* adrs, byte* node) + int z, const byte* pk_seed, word32* adrs, byte* node, + SlhDsaWotsBufs* bufs) { int ret = 0; @@ -5028,7 +5176,8 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, /* Step 3: Set key pair address. */ HA_SetKeyPairAddress(adrs, i); /* Step 4: Generate WOTS+ public key. */ - ret = slhdsakey_wots_pkgen(key, sk_seed, pk_seed, adrs, node); + ret = slhdsakey_wots_pkgen(key, sk_seed, pk_seed, adrs, node, + bufs); } else { WC_DECLARE_VAR(nodes, byte, (SLHDSA_MAX_H_M + 2) * SLHDSA_MAX_N, @@ -5049,7 +5198,7 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, HA_SetKeyPairAddress(adrs, m * (word32)i + j); /* Step 4: Generate WOTS+ public key. */ ret = slhdsakey_wots_pkgen(key, sk_seed, pk_seed, adrs, - nodes + ((word32)z - 1U + (j & 1U)) * n); + nodes + ((word32)z - 1U + (j & 1U)) * n, bufs); if (ret != 0) { break; } @@ -5126,7 +5275,8 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, - int z, const byte* pk_seed, word32* adrs, byte* node) + int z, const byte* pk_seed, word32* adrs, byte* node, + SlhDsaWotsBufs* bufs) { int ret; byte nodes[2 * SLHDSA_MAX_N]; @@ -5138,18 +5288,19 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, /* Step 3: Set key pair address. */ HA_SetKeyPairAddress(adrs, i); /* Step 4: Generate WOTS+ public key. */ - ret = slhdsakey_wots_pkgen(key, sk_seed, pk_seed, adrs, node); + ret = slhdsakey_wots_pkgen(key, sk_seed, pk_seed, adrs, node, + bufs); } else { byte n = key->params->n; /* Step 6: Calculate left node recursively. */ ret = slhdsakey_xmss_node(key, sk_seed, 2 * i, z - 1, pk_seed, adrs, - nodes); + nodes, bufs); if (ret == 0) { /* Step 7: Calculate right node recursively. */ ret = slhdsakey_xmss_node(key, sk_seed, 2 * i + 1, z - 1, pk_seed, - adrs, nodes + n); + adrs, nodes + n, bufs); } if (ret == 0) { /* Steps 8-10: Step type, height and index for TREE. */ @@ -5191,11 +5342,95 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ +/* Allocate the buffers a subtree's WOTS+ public keys share. + * + * @param [in] key SLH-DSA key. + * @param [out] bufs Buffers to allocate. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ +static int slhdsakey_wots_bufs_alloc(SlhDsaKey* key, SlhDsaWotsBufs* bufs) +{ + int ret = 0; + + bufs->sk = (byte*)XMALLOC(SLHDSA_WOTS_SK_SZ, key->heap, + DYNAMIC_TYPE_SLHDSA); + if (bufs->sk == NULL) { + ret = MEMORY_E; + } +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + bufs->fixed = NULL; + bufs->state = NULL; + if (ret == 0) { + bufs->fixed = (word64*)XMALLOC(sizeof(word64) * SLHDSA_WOTS_FIXED_W, + key->heap, DYNAMIC_TYPE_SLHDSA); + if (bufs->fixed == NULL) { + ret = MEMORY_E; + } + } + if (ret == 0) { + bufs->state = (word64*)XMALLOC(sizeof(word64) * SLHDSA_SHAKE_STATE_W, + key->heap, DYNAMIC_TYPE_SLHDSA); + if (bufs->state == NULL) { + ret = MEMORY_E; + } + } +#endif + + return ret; +} + +/* Release the buffers a subtree's WOTS+ public keys share. + * + * @param [in] key SLH-DSA key. + * @param [in, out] bufs Buffers to release. + */ +static void slhdsakey_wots_bufs_free(SlhDsaKey* key, SlhDsaWotsBufs* bufs) +{ +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + XFREE(bufs->state, key->heap, DYNAMIC_TYPE_SLHDSA); + XFREE(bufs->fixed, key->heap, DYNAMIC_TYPE_SLHDSA); +#endif + XFREE(bufs->sk, key->heap, DYNAMIC_TYPE_SLHDSA); +} + +/* Compute the root of the top XMSS subtree from the private key seeds. + * + * FIPS 205. Section 9.1. Algorithm 18. Steps 1-3. + * + * @param [in] key SLH-DSA key with the private key seeds set. + * @param [out] root Root node. n bytes. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + * @return SHAKE-256 error return code on digest failure. + */ +static int slhdsakey_root_from_seed(SlhDsaKey* key, byte* root) +{ + int ret; + byte n = key->params->n; + HashAddress adrs; + /* Buffers every WOTS+ public key in the root subtree shares. */ + SlhDsaWotsBufs bufs; + + ret = slhdsakey_wots_bufs_alloc(key, &bufs); + if (ret == 0) { + /* Steps 1-2: Address of the top-level XMSS tree. */ + HA_Init(adrs); + HA_SetLayerAddress(adrs, key->params->d - 1); + /* Step 3: Compute the root node. */ + ret = slhdsakey_xmss_node(key, key->sk, 0, key->params->h_m, + key->sk + 2 * n, adrs, root, &bufs); + } + slhdsakey_wots_bufs_free(key, &bufs); + + return ret; +} + static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, const byte* sk_seed, word32 idx, const byte* pk_seed, word32* adrs, byte* sig_xmss) { - int ret = WC_NO_ERR_TRACE(BAD_FUNC_ARG); + int ret = 0; byte n = key->params->n; byte len = key->params->len; byte h_m = key->params->h_m; @@ -5203,14 +5438,37 @@ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, byte* auth = sig_xmss + (len * n); word32 i = idx; int j; + /* Buffers every WOTS+ public key in this subtree shares. */ + SlhDsaWotsBufs bufs; + WC_DECLARE_VAR(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap); +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + WC_DECLARE_VAR(fixed, word64, SLHDSA_WOTS_FIXED_W, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_STATE_W, key->heap); +#endif + + WC_ALLOC_VAR_EX(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_WOTS_FIXED_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + bufs.fixed = fixed; + bufs.state = state; +#endif + bufs.sk = sk; /* Step 1: For each height of XMSS tree. */ - for (j = 0; j < h_m; j++) { + for (j = 0; (ret == 0) && (j < h_m); j++) { /* Step 2: Calculate index of other node. */ word32 k = i ^ 1; /* Step 3: Calculate authentication node. */ ret = slhdsakey_xmss_node(key, sk_seed, (int)k, j, pk_seed, adrs, - auth); + auth, &bufs); if (ret != 0) { break; } @@ -5229,6 +5487,11 @@ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, ret = slhdsakey_wots_sign(key, m, sk_seed, pk_seed, adrs, sig_xmss); } +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); +#endif + WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } #endif /* !WOLFSSL_SLHDSA_VERIFY_ONLY */ @@ -7671,9 +7934,8 @@ int wc_SlhDsaKey_MakeKey(SlhDsaKey* key, WC_RNG* rng) */ static int slhdsakey_compute_root(SlhDsaKey* key) { - int ret = 0; - byte n = key->params->n; - HashAddress adrs; + int ret = 0; + byte n = key->params->n; #ifdef WOLFSSL_SLHDSA_SHA2 /* Pre-compute SHA2 midstates now that PK.seed is set. */ @@ -7685,13 +7947,8 @@ static int slhdsakey_compute_root(SlhDsaKey* key) } #endif - /* Step 1: Set address to all zeroes. */ - HA_Init(adrs); - /* Step 2: Set the address layer to the top of the subtree. */ - HA_SetLayerAddress(adrs, key->params->d - 1); - /* Step 3: Compute the root node. */ - ret = slhdsakey_xmss_node(key, key->sk, 0, key->params->h_m, - key->sk + 2 * n, adrs, &key->sk[3 * n]); + /* Steps 1-3: Compute the root node from the seeds now in the key. */ + ret = slhdsakey_root_from_seed(key, &key->sk[3 * n]); if (ret == 0) { key->flags = WC_SLHDSA_FLAG_BOTH_KEYS; } @@ -7809,7 +8066,6 @@ int wc_SlhDsaKey_MakeKeyWithRandom(SlhDsaKey* key, const byte* sk_seed, byte pct_root[SLHDSA_MAX_N]; byte pct_pub[2 * SLHDSA_MAX_N]; word32 pct_pubLen = (word32)(n * 2); - HashAddress pct_adrs; /* The key identifier the IG names. */ ret = wc_SlhDsaKey_ExportPublic(key, pct_pub, &pct_pubLen); @@ -7819,10 +8075,7 @@ int wc_SlhDsaKey_MakeKeyWithRandom(SlhDsaKey* key, const byte* sk_seed, /* The public root, recomputed from the private seed. */ if (ret == 0) { - HA_Init(pct_adrs); - HA_SetLayerAddress(pct_adrs, key->params->d - 1); - ret = slhdsakey_xmss_node(key, key->sk, 0, key->params->h_m, - key->sk + 2 * n, pct_adrs, pct_root); + ret = slhdsakey_root_from_seed(key, pct_root); if ((ret == 0) && (XMEMCMP(pct_root, key->sk + 3 * n, n) != 0)) { ret = SLH_DSA_PCT_E; } @@ -9830,18 +10083,14 @@ int wc_SlhDsaKey_CheckKey(SlhDsaKey* key) } #else if (ret == 0) { - byte n = key->params->n; - byte root[SLHDSA_MAX_N]; - HashAddress adrs; + byte n = key->params->n; + byte root[SLHDSA_MAX_N]; /* Recompute the public root from the private seed and compare. * Done directly rather than by regenerating the key: regeneration * overwrites the key being checked, and its key-pair test frees the * key on failure, which a validation call must never do. */ - HA_Init(adrs); - HA_SetLayerAddress(adrs, key->params->d - 1); - ret = slhdsakey_xmss_node(key, key->sk, 0, key->params->h_m, - key->sk + 2 * n, adrs, root); + ret = slhdsakey_root_from_seed(key, root); if ((ret == 0) && (XMEMCMP(root, key->sk + 3 * n, n) != 0)) { ret = WC_KEY_MISMATCH_E; } From d210652ce2de33c151862b6292c6fad6e8805988 Mon Sep 17 00:00:00 2001 From: Tobias Frauenschlaeger Date: Wed, 19 Aug 2026 11:30:39 +0000 Subject: [PATCH 5/6] Fix SLH-DSA memory optimization review findings The eight-way Keccak permutation is emitted by the SHA-3 assembly only for ML-KEM and ML-DSA builds, so an SLH-DSA build with the Intel assembly and neither of those failed to link. Add WOLFSSL_HAVE_SLHDSA to the x8 guards, matching the x4 ones, and close the guard after the permutation so the seed helpers that follow stay as they were. Two entries of the WOLFSSL_HAVE_SHA256_HASH_BLOCK port list named macros a port only sets inside its own .c file, so they read as undefined in the header and excluded nothing. Test the gate macros instead, and record the constraint. Keygen allocated the shared WOTS+ buffers with XMALLOC while xmss_sign used WC_DECLARE_VAR. Consolidate on the latter: heap-free builds work again, and a partial allocation no longer leaks. Guard the chain value buffer so a small memory build without the batched path stops carrying one it never reads. Clear the Keccak state and node buffers that hold FORS secret keys before releasing them, and clear the SHA-256 block helper's working buffers unconditionally rather than per call site. Also: reject a message that cannot fit the block with its padding, return early from a zero step chain, assert the len parity the buffer overshoot depends on, and correct the parameter documentation the refactor left behind. Added two CI configs covering the build breaks: SLH-DSA with the assembly but without ML-KEM/ML-DSA, and LMS/SLH-DSA without raw hash access. --- .github/configs/pq-all.json | 14 ++ wolfcrypt/src/sha3_asm.S | 8 +- wolfcrypt/src/sha3_asm.asm | 15 ++ wolfcrypt/src/wc_slhdsa.c | 390 ++++++++++++++++++++---------------- wolfssl/wolfcrypt/sha256.h | 24 ++- 5 files changed, 266 insertions(+), 185 deletions(-) diff --git a/.github/configs/pq-all.json b/.github/configs/pq-all.json index a250cadbb64..de4a8915773 100644 --- a/.github/configs/pq-all.json +++ b/.github/configs/pq-all.json @@ -154,6 +154,20 @@ "--enable-dtls-mtu", "--enable-dtls-frag-ch", "--enable-dtlscid", "--enable-mlkem=make,enc,dec,768", "--disable-qt", "CPPFLAGS=-pedantic -Wdeclaration-after-statement -Wnull-dereference -DWOLFCRYPT_TEST_LINT -DNO_WOLFSSL_CIPHER_SUITE_TEST -DTEST_LIBWOLFSSL_SOURCES_INCLUSION_SEQUENCE"]}, +{"name": "slhdsa-only-asm", "minutes": 3, + "comment": "SLH-DSA with the Intel assembly but without ML-KEM/ML-DSA; the 8-way AVX512 Keccak permutation is emitted for those, so this catches SLH-DSA calling a symbol the SHA-3 assembly did not build", + "configure": ["--enable-intelasm", "--enable-sp-asm", + "--enable-slhdsa=yes,sha2", "--disable-mlkem", "--disable-mldsa", + "--disable-dilithium"]}, +{"name": "slhdsa-no-shake-x8", "minutes": 3, + "comment": "SLH-DSA with the Intel assembly but the eight-way AVX512 Keccak path pinned off, so the four-way path is covered whatever the runner CPU; paired with slhdsa-only-asm, a KAT that passes here and fails there isolates the eight-way code", + "configure": ["--enable-intelasm", "--enable-sp-asm", "--enable-slhdsa", + "CPPFLAGS=-DWOLFSSL_SLHDSA_NO_SHAKE_X8"]}, +{"name": "lms-slhdsa-no-hash-raw", "minutes": 3, + "comment": "No raw hash access, as the PSA and hardware SHA-256 ports have; wc_Sha256HashBlock() is not built, so this catches LMS, SLH-DSA or the test suite calling it anyway", + "configure": ["--enable-lms", "--enable-xmss", + "--enable-slhdsa=yes,sha2", + "CPPFLAGS=-DWOLFSSL_NO_HASH_RAW"]}, {"name": "mldsa-no-asn1-opensslextra", "minutes": 2.8, "configure": ["--enable-intelasm", "--enable-sp-asm", "--enable-dilithium=yes", "--enable-opensslextra", diff --git a/wolfcrypt/src/sha3_asm.S b/wolfcrypt/src/sha3_asm.S index 8eb74805876..424db5fbf8a 100644 --- a/wolfcrypt/src/sha3_asm.S +++ b/wolfcrypt/src/sha3_asm.S @@ -36870,7 +36870,7 @@ _sha3_256_blocksx4_seed_64_avx2: #endif /* HAVE_INTEL_AVX512 */ #endif /* NO_AVX512_SUPPORT */ #ifdef HAVE_INTEL_AVX512 -#if defined(WOLFSSL_HAVE_FRODOKEM) || defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) +#if defined(WOLFSSL_HAVE_FRODOKEM) || defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) || defined(WOLFSSL_HAVE_SLHDSA) #ifndef __APPLE__ .data #else @@ -36978,7 +36978,7 @@ L_sha3_x8_avx512_r: .quad 0x8000000080008008,0x8000000080008008 .quad 0x8000000080008008,0x8000000080008008 .quad 0x8000000080008008,0x8000000080008008 -#endif /* defined(WOLFSSL_HAVE_FRODOKEM) || defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) */ +#endif /* defined(WOLFSSL_HAVE_FRODOKEM) || defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) || defined(WOLFSSL_HAVE_SLHDSA) */ #ifdef WOLFSSL_HAVE_FRODOKEM #ifndef __APPLE__ .text @@ -39922,7 +39922,7 @@ L_sha3_blocksx8_out_avx512_done: .size sha3_blocksx8_out_avx512,.-sha3_blocksx8_out_avx512 #endif /* __APPLE__ */ #endif /* WOLFSSL_HAVE_FRODOKEM */ -#if defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) +#if defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) || defined(WOLFSSL_HAVE_SLHDSA) #ifndef __APPLE__ .text .globl sha3_blocksx8_avx512 @@ -42801,6 +42801,8 @@ _sha3_blocksx8_avx512: #ifndef __APPLE__ .size sha3_blocksx8_avx512,.-sha3_blocksx8_avx512 #endif /* __APPLE__ */ +#endif /* defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) || defined(WOLFSSL_HAVE_SLHDSA) */ +#if defined(WOLFSSL_HAVE_MLKEM) || defined(WOLFSSL_HAVE_MLDSA) #ifndef __APPLE__ .data #else diff --git a/wolfcrypt/src/sha3_asm.asm b/wolfcrypt/src/sha3_asm.asm index 1f023db1ed7..fe8178c2150 100644 --- a/wolfcrypt/src/sha3_asm.asm +++ b/wolfcrypt/src/sha3_asm.asm @@ -36889,6 +36889,9 @@ ENDIF IFDEF WOLFSSL_HAVE_MLDSA wc_masm_cond_2 = 1 ENDIF +IFDEF WOLFSSL_HAVE_SLHDSA +wc_masm_cond_2 = 1 +ENDIF IF wc_masm_cond_2 _DATA SEGMENT ALIGN 16 @@ -39954,6 +39957,9 @@ ENDIF IFDEF WOLFSSL_HAVE_MLDSA wc_masm_cond_3 = 1 ENDIF +IFDEF WOLFSSL_HAVE_SLHDSA +wc_masm_cond_3 = 1 +ENDIF IF wc_masm_cond_3 _TEXT SEGMENT READONLY PARA sha3_blocksx8_avx512 PROC @@ -42842,6 +42848,15 @@ sha3_blocksx8_avx512 PROC ret sha3_blocksx8_avx512 ENDP _TEXT ENDS +ENDIF +wc_masm_cond_3b = 0 +IFDEF WOLFSSL_HAVE_MLKEM +wc_masm_cond_3b = 1 +ENDIF +IFDEF WOLFSSL_HAVE_MLDSA +wc_masm_cond_3b = 1 +ENDIF +IF wc_masm_cond_3b _DATA SEGMENT ALIGN 16 L_sha3_128_blocksx8_seed_avx512_end_mark QWORD 8000000000000000h, 8000000000000000h diff --git a/wolfcrypt/src/wc_slhdsa.c b/wolfcrypt/src/wc_slhdsa.c index fcae7ff31a1..edabb8462df 100644 --- a/wolfcrypt/src/wc_slhdsa.c +++ b/wolfcrypt/src/wc_slhdsa.c @@ -212,25 +212,39 @@ wc_static_assert(SLHDSA_MAX_MSG_SZ <= 255); wc_static_assert(SLHDSA_WOTS_SK_SZ >= (((SLHDSA_MAX_MSG_SZ + 7) / 8) * 8) * SLHDSA_MAX_N); -/* Words in the unchanging head of the batched state, widest path in build. */ +/* Words in the unchanging head of the batched state: what + * slhdsakey_chain_x8() and slhdsakey_chain_x4_32() write before the hashes. */ #ifdef SLHDSA_HAVE_SHAKE_X8 #define SLHDSA_WOTS_FIXED_W SLHDSA_SHAKE_X8_FIXED_W #else #define SLHDSA_WOTS_FIXED_W (8 * 4) #endif -/* Buffers reused by every WOTS+ public key computed from one subtree. A - * subtree builds thousands of them, so allocating per public key put the - * allocator on the hot path. */ +/* Only the batched WOTS+ public key paths and the full arm of + * slhdsakey_wots_pkgen_chain_c() read the shared chain value buffer. */ +#if !defined(WOLFSSL_WC_SLHDSA_SMALL_MEM) || \ + (defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL)) + #define SLHDSA_NEED_WOTS_SK_BUF +#endif + +/* Buffers shared by every WOTS+ public key of one subtree. */ typedef struct SlhDsaWotsBufs { - /* Chain values of the public key being built. */ +#ifdef SLHDSA_NEED_WOTS_SK_BUF + /* Chain values of the public key being built. Non-NULL by the time a chain + * helper runs: both owners abandon the subtree when an allocation fails. */ byte* sk; +#endif #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) /* Unchanging head of the batched Keccak state. */ word64* fixed; /* Batched Keccak state. */ word64* state; #endif +#if !defined(SLHDSA_NEED_WOTS_SK_BUF) && \ + !(defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL)) + /* C89 has no empty struct, and this build shares no buffers. */ + byte unused; +#endif } SlhDsaWotsBufs; #ifndef WC_SLHDSA_ALL_NO_256F @@ -766,7 +780,10 @@ static int slhdsakey_hash_shake_4(wc_Shake* shake, const byte* data1, #ifdef SLHDSA_SHA2_BLOCK_HASH /* A registered callback expects to see every update, so the direct path is - * only taken for an object no callback has claimed. */ + * only taken for an object no callback has claimed. The key's own hash objects + * are always initialised with INVALID_DEVID, so today this is constant true + * and the direct path always runs; the check is what stops a later change to + * that initialisation from silently bypassing a registered callback. */ #ifdef WOLF_CRYPTO_CB #define SLHDSA_SHA256_RAW_OK(key) \ ((key)->hash.sha2.sha256.devId == INVALID_DEVID) @@ -871,8 +888,10 @@ wc_static_assert(SLHDSA_HAC_SZ + 32 + 1 + 8 <= WC_SHA256_BLOCK_SIZE); * @param [in] m2_len Length of second message part. * @param [out] hash Buffer to hold hash output. * @param [in] hash_len Number of bytes of hash to output. - * @param [in] zeroize Wipe the working buffers when the input is secret. + * @param [in] zeroize Track the working buffers as secret for the memory + * zeroization check. They are wiped either way. * @return 0 on success. + * @return BUFFER_E when the message does not fit the block with its padding. */ static int slhdsakey_sha256_block_hash(SlhDsaKey* key, const byte* address, const byte* m1, byte m1_len, const byte* m2, byte m2_len, byte* hash, @@ -887,6 +906,11 @@ static int slhdsakey_sha256_block_hash(SlhDsaKey* key, const byte* address, /* Length covers the midstate block as well as this one. */ word32 bits = (WC_SHA256_BLOCK_SIZE + len) * 8; + /* Message, padding byte and 8 length bytes have to share one block. */ + if (len + 1 + 8 > WC_SHA256_BLOCK_SIZE) { + return BUFFER_E; + } + XMEMCPY(block, address, SLHDSA_HAC_SZ); XMEMCPY(block + SLHDSA_HAC_SZ, m1, m1_len); if (m2_len > 0) { @@ -910,6 +934,7 @@ static int slhdsakey_sha256_block_hash(SlhDsaKey* key, const byte* address, wc_MemZero_Add("slhdsa sha256 block", block, WC_SHA256_BLOCK_SIZE); wc_MemZero_Add("slhdsa sha256 digest", digest, sizeof(digest)); #endif + /* Cleared unconditionally rather than per call site. */ ForceZero(block, WC_SHA256_BLOCK_SIZE); ForceZero(digest, sizeof(digest)); #ifdef WOLFSSL_CHECK_MEM_ZERO @@ -2709,6 +2734,11 @@ static int slhdsakey_chain_idx_x8(byte* sk, word32 i, word32 s, byte n, /* Words the eight hashes occupy in the state. */ word32 hw = (word32)(n / 8) * 8; + /* chain() over zero steps returns its input unchanged. */ + if (s == 0) { + return 0; + } + slhdsakey_shake256_set_hash_x8(state, o, sk, n); for (j = i; j < i + s; j++) { @@ -2753,9 +2783,8 @@ static int slhdsakey_chain_idx_x8(byte* sk, word32 i, word32 s, byte n, * @param [in] n Number of bytes in hash output. * @param [in] ca Chain address start index. * @param [out] sk Buffer to hold hash output. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_hash_prf_x4(const byte* pk_seed, const byte* sk_seed, @@ -2796,9 +2825,9 @@ static int slhdsakey_hash_prf_x4(const byte* pk_seed, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] fixed Caller owned buffer for the unchanging state head. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x4_16(byte* sk, const byte* pk_seed, byte* addr, byte ca, word64* fixed, word64* state) @@ -2850,9 +2879,9 @@ static int slhdsakey_chain_x4_16(byte* sk, const byte* pk_seed, byte* addr, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] fixed Caller owned buffer for the unchanging state head. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x4_24(byte* sk, const byte* pk_seed, byte* addr, byte ca, word64* fixed, word64* state) @@ -2904,9 +2933,9 @@ static int slhdsakey_chain_x4_24(byte* sk, const byte* pk_seed, byte* addr, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] fixed Caller owned buffer for the unchanging state head. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_chain_x4_32(byte* sk, const byte* pk_seed, byte* addr, byte ca, word64* fixed, word64* state) @@ -3265,8 +3294,9 @@ static int slhdsakey_chain_idx_32(SlhDsaKey* key, byte* sk, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] sk_addr Encoded WOTS PRF HashAddress. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, const byte* pk_seed, byte* addr, byte* sk_addr, @@ -3337,8 +3367,9 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] sk_addr Encoded WOTS PRF HashAddress. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, const byte* pk_seed, byte* addr, byte* sk_addr, @@ -3409,8 +3440,9 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] sk_addr Encoded WOTS PRF HashAddress. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, const byte* pk_seed, byte* addr, byte* sk_addr, @@ -3624,15 +3656,16 @@ static int slhdsakey_chain_x8(byte* sk, const byte* pk_seed, byte* addr, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded WOTS HASH HashAddress. * @param [in] sk_addr Encoded WOTS PRF HashAddress. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, const byte* pk_seed, byte* addr, byte* sk_addr, SlhDsaWotsBufs* bufs) { int ret = 0; - int i = 0; + int i; byte n = key->params->n; byte len = key->params->len; byte* sk = bufs->sk; @@ -3696,8 +3729,9 @@ static int slhdsakey_wots_pkgen_chain_x8(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] adrs HashAddress. * @param [in] sk_adrs WOTS PRF HashAddress. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, const byte* pk_seed, word32* adrs, word32* sk_adrs, SlhDsaWotsBufs* bufs) @@ -3712,7 +3746,7 @@ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, HA_Encode(adrs, addr); #ifdef SLHDSA_HAVE_SHAKE_X8 - if (USE_INTEL_AVX512(cpuid_flags)) { + if (SLHDSA_USE_SHAKE_X8()) { return slhdsakey_wots_pkgen_chain_x8(key, sk_seed, pk_seed, addr, sk_addr, bufs); } @@ -3770,8 +3804,9 @@ static int slhdsakey_wots_pkgen_chain_x4(SlhDsaKey* key, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] adrs HashAddress. * @param [in] sk_adrs WOTS PRF HashAddress. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, @@ -3786,26 +3821,23 @@ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, #if !defined(WOLFSSL_WC_SLHDSA_SMALL_MEM) byte* sk = bufs->sk; - XMEMSET(sk, 0, SLHDSA_WOTS_SK_SZ); - if (ret == 0) { - /* Step 4. len consecutive addresses. */ - for (i = 0; i < len; i++) { - /* Step 5. Set chain address for WOTS PRF. */ - HA_SetChainAddress(sk_adrs, i); - /* Step 6. PRF hash seeds and chain address. */ - ret = HASH_PRF(key, pk_seed, sk_seed, sk_adrs, n, - sk + i * n); - if (ret != 0) { - break; - } - /* Step 7. Set chain address for WOTS HASH. */ - HA_SetChainAddress(adrs, i); - /* Step 8. Chain hashes for w-1 iterations. */ - ret = slhdsakey_chain(key, sk + i * n, 0, SLHDSA_WM1, pk_seed, adrs, - sk + i * n); - if (ret != 0) { - break; - } + XMEMSET(sk, 0, (word32)len * n); + /* Step 4. len consecutive addresses. */ + for (i = 0; i < len; i++) { + /* Step 5. Set chain address for WOTS PRF. */ + HA_SetChainAddress(sk_adrs, i); + /* Step 6. PRF hash seeds and chain address. */ + ret = HASH_PRF(key, pk_seed, sk_seed, sk_adrs, n, sk + i * n); + if (ret != 0) { + break; + } + /* Step 7. Set chain address for WOTS HASH. */ + HA_SetChainAddress(adrs, i); + /* Step 8. Chain hashes for w-1 iterations. */ + ret = slhdsakey_chain(key, sk + i * n, 0, SLHDSA_WM1, pk_seed, adrs, + sk + i * n); + if (ret != 0) { + break; } } if (ret == 0) { @@ -3875,13 +3907,14 @@ static int slhdsakey_wots_pkgen_chain_c(SlhDsaKey* key, const byte* sk_seed, * 13: pk <- Tlen(PK.seed, wotspkADRS, tmp) > compress public key * 14: return pk * - * @param [in] key SLH-DSA key. - * @param [in] sk_seed Private key seed. - * @param [in] pk_seed Public key seed. - * @param [in] adrs HashAddress. - * @param [in] sk_adrs WOTS PRF HashAddress. + * @param [in] key SLH-DSA key. + * @param [in] sk_seed Private key seed. + * @param [in] pk_seed Public key seed. + * @param [in] adrs HashAddress. + * @param [out] node WOTS+ public key. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_wots_pkgen(SlhDsaKey* key, const byte* sk_seed, @@ -4625,32 +4658,6 @@ static int slhdsakey_chain_idx_to_max_32(SlhDsaKey* key, const byte* sig, #endif #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) -/* Computes a WOTS+ public key from a message and its signature. - * - * Computes four iteration hashes simultaneously. - * - * FIPS 205. Section 5.3. Algorithm 8. - * wots_pkFromSig(sig, M, PK.seed, ADRS) - * ... - * 8: for i from 0 to len - 1 do - * 9: ADRS.setChainAddress(i) - * ... - * 11: end for - * 12: wotspkADRS <- ADRS > copy address to create WOTS+ public key address - * 13: wotspkADRS.setTypeAndClear(WOTS_PK) - * 14: wotspkADRS.setKeyPairAddress(ADRS.getKeyPairAddress()) - * 15: pksig <- Tlen (PK.seed, wotspkADRS, tmp) - * 16: return pksig - * - * @param [in] key SLH-DSA key. - * @param [in] sig Signature - (2.n + 3) hashes of length n. - * @param [in] msg Encoded message with checksum. - * @param [in] pk_seed Public key seed. - * @param [in] adrs WOTS HASH HashAddress. - * @param [out] pk_sig Root node - public key signature. - * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. - */ #ifdef SLHDSA_HAVE_SHAKE_X8 /* Complete the WOTS+ chains of a signature, eight at a time. * @@ -4768,6 +4775,32 @@ static int slhdsakey_wots_pk_from_sig_x8(SlhDsaKey* key, const byte* sig, } #endif /* SLHDSA_HAVE_SHAKE_X8 */ +/* Computes a WOTS+ public key from a message and its signature. + * + * Computes four iteration hashes simultaneously. + * + * FIPS 205. Section 5.3. Algorithm 8. + * wots_pkFromSig(sig, M, PK.seed, ADRS) + * ... + * 8: for i from 0 to len - 1 do + * 9: ADRS.setChainAddress(i) + * ... + * 11: end for + * 12: wotspkADRS <- ADRS > copy address to create WOTS+ public key address + * 13: wotspkADRS.setTypeAndClear(WOTS_PK) + * 14: wotspkADRS.setKeyPairAddress(ADRS.getKeyPairAddress()) + * 15: pksig <- Tlen (PK.seed, wotspkADRS, tmp) + * 16: return pksig + * + * @param [in] key SLH-DSA key. + * @param [in] sig Signature - (2.n + 3) hashes of length n. + * @param [in] msg Encoded message with checksum. + * @param [in] pk_seed Public key seed. + * @param [in] adrs WOTS HASH HashAddress. + * @param [out] pk_sig Root node - public key signature. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + */ static int slhdsakey_wots_pk_from_sig_x4(SlhDsaKey* key, const byte* sig, const byte* msg, const byte* pk_seed, word32* adrs, byte* pk_sig) { @@ -5159,8 +5192,9 @@ static int slhdsakey_wots_pk_from_sig(SlhDsaKey* key, const byte* sig, * @param [in] pk_seed Public key seed. * @param [in, out] adrs HashAddress - WOTS HASH. * @param [out] node Root node. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, @@ -5270,8 +5304,9 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, * @param [in] pk_seed Public key seed. * @param [in, out] adrs HashAddress - WOTS HASH. * @param [out] node Root node. + * @param [in] bufs Buffers shared by the WOTS+ public keys of one + * subtree. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, @@ -5316,84 +5351,6 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, } #endif -/* Generate XMSS signature. - * - * FIPS 205. Section 6.2. Algorithm 10. - * xmss_sign(M SK.seed, idx PK.seed, ADRS) - * 1: for j from 0 to h' - 1 do > build authentication path - * 2: k <- lower(idx/2^j) XOR 1 - * 3: AUTH[j] <- xmss_node(SK.seed, k, j, PK.seed, ADRS) - * 4: end for - * 5: ADRS.setTypeAndClear(WOTS_HASH) - * 6: ADRS.setKeyPairAddress(idx) - * 7: sig <- wots_sign(M , SK.seed, PK.seed, ADRS) - * 8: SIGXMSS <- sig || AUTH - * 9: return SIGXMSS - * - * @param [in] key SLH-DSA key. - * @param [in] m n-byte message. - * @param [in] sk_seed Private key seed. - * @param [in] idx Key pair address of WOTS hash. - * @param [in] pk_seed Public key seed. - * @param [in] adrs HashAddress. - * @param [out] sig_xmss XMSS signature. - * len n-byte nodes and h' authentication nodes. - * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. - * @return SHAKE-256 error return code on digest failure. - */ -/* Allocate the buffers a subtree's WOTS+ public keys share. - * - * @param [in] key SLH-DSA key. - * @param [out] bufs Buffers to allocate. - * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. - */ -static int slhdsakey_wots_bufs_alloc(SlhDsaKey* key, SlhDsaWotsBufs* bufs) -{ - int ret = 0; - - bufs->sk = (byte*)XMALLOC(SLHDSA_WOTS_SK_SZ, key->heap, - DYNAMIC_TYPE_SLHDSA); - if (bufs->sk == NULL) { - ret = MEMORY_E; - } -#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) - bufs->fixed = NULL; - bufs->state = NULL; - if (ret == 0) { - bufs->fixed = (word64*)XMALLOC(sizeof(word64) * SLHDSA_WOTS_FIXED_W, - key->heap, DYNAMIC_TYPE_SLHDSA); - if (bufs->fixed == NULL) { - ret = MEMORY_E; - } - } - if (ret == 0) { - bufs->state = (word64*)XMALLOC(sizeof(word64) * SLHDSA_SHAKE_STATE_W, - key->heap, DYNAMIC_TYPE_SLHDSA); - if (bufs->state == NULL) { - ret = MEMORY_E; - } - } -#endif - - return ret; -} - -/* Release the buffers a subtree's WOTS+ public keys share. - * - * @param [in] key SLH-DSA key. - * @param [in, out] bufs Buffers to release. - */ -static void slhdsakey_wots_bufs_free(SlhDsaKey* key, SlhDsaWotsBufs* bufs) -{ -#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) - XFREE(bufs->state, key->heap, DYNAMIC_TYPE_SLHDSA); - XFREE(bufs->fixed, key->heap, DYNAMIC_TYPE_SLHDSA); -#endif - XFREE(bufs->sk, key->heap, DYNAMIC_TYPE_SLHDSA); -} - /* Compute the root of the top XMSS subtree from the private key seeds. * * FIPS 205. Section 9.1. Algorithm 18. Steps 1-3. @@ -5406,13 +5363,37 @@ static void slhdsakey_wots_bufs_free(SlhDsaKey* key, SlhDsaWotsBufs* bufs) */ static int slhdsakey_root_from_seed(SlhDsaKey* key, byte* root) { - int ret; + int ret = 0; byte n = key->params->n; HashAddress adrs; /* Buffers every WOTS+ public key in the root subtree shares. */ SlhDsaWotsBufs bufs; +#ifdef SLHDSA_NEED_WOTS_SK_BUF + WC_DECLARE_VAR(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap); +#endif +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + WC_DECLARE_VAR(fixed, word64, SLHDSA_WOTS_FIXED_W, key->heap); + WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_STATE_W, key->heap); +#endif + +#ifdef SLHDSA_NEED_WOTS_SK_BUF + WC_ALLOC_VAR_EX(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + bufs.sk = sk; +#endif +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + if (ret == 0) { + WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_WOTS_FIXED_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + if (ret == 0) { + WC_ALLOC_VAR_EX(state, word64, SLHDSA_SHAKE_STATE_W, key->heap, + DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + } + bufs.fixed = fixed; + bufs.state = state; +#endif - ret = slhdsakey_wots_bufs_alloc(key, &bufs); if (ret == 0) { /* Steps 1-2: Address of the top-level XMSS tree. */ HA_Init(adrs); @@ -5421,11 +5402,43 @@ static int slhdsakey_root_from_seed(SlhDsaKey* key, byte* root) ret = slhdsakey_xmss_node(key, key->sk, 0, key->params->h_m, key->sk + 2 * n, adrs, root, &bufs); } - slhdsakey_wots_bufs_free(key, &bufs); +#if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) + WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); + WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); +#endif +#ifdef SLHDSA_NEED_WOTS_SK_BUF + WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); +#endif return ret; } +/* Generate XMSS signature. + * + * FIPS 205. Section 6.2. Algorithm 10. + * xmss_sign(M SK.seed, idx PK.seed, ADRS) + * 1: for j from 0 to h' - 1 do > build authentication path + * 2: k <- lower(idx/2^j) XOR 1 + * 3: AUTH[j] <- xmss_node(SK.seed, k, j, PK.seed, ADRS) + * 4: end for + * 5: ADRS.setTypeAndClear(WOTS_HASH) + * 6: ADRS.setKeyPairAddress(idx) + * 7: sig <- wots_sign(M , SK.seed, PK.seed, ADRS) + * 8: SIGXMSS <- sig || AUTH + * 9: return SIGXMSS + * + * @param [in] key SLH-DSA key. + * @param [in] m n-byte message. + * @param [in] sk_seed Private key seed. + * @param [in] idx Key pair address of WOTS hash. + * @param [in] pk_seed Public key seed. + * @param [in] adrs HashAddress. + * @param [out] sig_xmss XMSS signature. + * len n-byte nodes and h' authentication nodes. + * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. + * @return SHAKE-256 error return code on digest failure. + */ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, const byte* sk_seed, word32 idx, const byte* pk_seed, word32* adrs, byte* sig_xmss) @@ -5440,14 +5453,19 @@ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, int j; /* Buffers every WOTS+ public key in this subtree shares. */ SlhDsaWotsBufs bufs; +#ifdef SLHDSA_NEED_WOTS_SK_BUF WC_DECLARE_VAR(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap); +#endif #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) WC_DECLARE_VAR(fixed, word64, SLHDSA_WOTS_FIXED_W, key->heap); WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_STATE_W, key->heap); #endif +#ifdef SLHDSA_NEED_WOTS_SK_BUF WC_ALLOC_VAR_EX(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); + bufs.sk = sk; +#endif #if defined(USE_INTEL_SPEEDUP) && !defined(WOLFSSL_WC_SLHDSA_SMALL) if (ret == 0) { WC_ALLOC_VAR_EX(fixed, word64, SLHDSA_WOTS_FIXED_W, key->heap, @@ -5460,7 +5478,6 @@ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, bufs.fixed = fixed; bufs.state = state; #endif - bufs.sk = sk; /* Step 1: For each height of XMSS tree. */ for (j = 0; (ret == 0) && (j < h_m); j++) { @@ -5491,7 +5508,9 @@ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); WC_FREE_VAR_EX(fixed, key->heap, DYNAMIC_TYPE_SLHDSA); #endif +#ifdef SLHDSA_NEED_WOTS_SK_BUF WC_FREE_VAR_EX(sk, key->heap, DYNAMIC_TYPE_SLHDSA); +#endif return ret; } #endif /* !WOLFSSL_SLHDSA_VERIFY_ONLY */ @@ -5840,9 +5859,8 @@ static int slhdsakey_fors_sk_gen(SlhDsaKey* key, const byte* sk_seed, * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. * @param [out] node Buffer to hold hash output. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_prf_ti_x4(const byte* pk_seed, const byte* sk_seed, byte* addr, byte n, word32 ti, byte* node, word64* state) @@ -5882,9 +5900,8 @@ static int slhdsakey_hash_prf_ti_x4(const byte* pk_seed, const byte* sk_seed, * @param [in, out] node On in, n-byte messages. On out, n-byte outputs. * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, byte n, word32 ti, word64* state) @@ -5910,6 +5927,9 @@ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, slhdsakey_shake256_get_hash_x4(state, node, n); } + /* state holds values derived from the FORS secret keys. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); + return ret; } @@ -5930,9 +5950,8 @@ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. * @param [out] hash Buffer to hold hash output. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_h_ti_x4(const byte* pk_seed, byte* addr, const byte* m, byte n, word32 ti, byte* hash, word64* state) @@ -5958,6 +5977,9 @@ static int slhdsakey_hash_h_ti_x4(const byte* pk_seed, byte* addr, slhdsakey_shake256_get_hash_x4(state, hash, n); } + /* state holds values derived from the FORS secret keys. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X4_STATE_W); + return ret; } @@ -5970,9 +5992,8 @@ static int slhdsakey_hash_h_ti_x4(const byte* pk_seed, byte* addr, * @param [in] n Number of bytes in each hash. * @param [in] ti Tree index start value. * @param [out] node Eight n-byte outputs. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_prf_ti_x8(const byte* pk_seed, const byte* sk_seed, byte* addr, byte n, word32 ti, byte* node, word64* state) @@ -6003,9 +6024,8 @@ static int slhdsakey_hash_prf_ti_x8(const byte* pk_seed, const byte* sk_seed, * @param [in, out] node On in, eight n-byte messages. On out, the outputs. * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_f_ti_x8(const byte* pk_seed, byte* addr, byte* node, byte n, word32 ti, word64* state) @@ -6031,6 +6051,9 @@ static int slhdsakey_hash_f_ti_x8(const byte* pk_seed, byte* addr, byte* node, slhdsakey_shake256_get_hash_x8(state, node, n); } + /* state holds values derived from the FORS secret keys. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + return ret; } @@ -6042,9 +6065,8 @@ static int slhdsakey_hash_f_ti_x8(const byte* pk_seed, byte* addr, byte* node, * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. * @param [out] hash Buffer to hold eight n-byte hash outputs. - * @param [in] heap Dynamic memory allocation hint. + * @param [in] state Caller owned Keccak state. * @return 0 on success. - * @return MEMORY_E on dynamic memory allocation failure. */ static int slhdsakey_hash_h_ti_x8(const byte* pk_seed, byte* addr, const byte* m, byte n, word32 ti, byte* hash, word64* state) @@ -6070,6 +6092,9 @@ static int slhdsakey_hash_h_ti_x8(const byte* pk_seed, byte* addr, slhdsakey_shake256_get_hash_x8(state, hash, n); } + /* state holds values derived from the FORS secret keys. */ + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_X8_STATE_W); + return ret; } #endif /* SLHDSA_HAVE_SHAKE_X8 */ @@ -6361,6 +6386,10 @@ static int slhdsakey_fors_node_x4_low(SlhDsaKey* key, const byte* sk_seed, ret = HASH_H(key, pk_seed, adrs, nodes, n, node); } + /* Holds FORS secret keys or values derived from them. */ + if (WC_VAR_OK(nodes)) { + ForceZero(nodes, (1 << SLHDSA_MAX_FORS_NODE_DEPTH) * SLHDSA_MAX_N); + } WC_FREE_VAR_EX(nodes, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } @@ -6485,6 +6514,10 @@ static int slhdsakey_fors_node_x4_high(SlhDsaKey* key, const byte* sk_seed, ret = HASH_H(key, pk_seed, adrs, nodes, n, node); } + /* Holds FORS secret keys or values derived from them. */ + if (WC_VAR_OK(nodes)) { + ForceZero(nodes, (1 << SLHDSA_MAX_FORS_NODE_TOP_DEPTH) * SLHDSA_MAX_N); + } WC_FREE_VAR_EX(nodes, key->heap, DYNAMIC_TYPE_SLHDSA); return ret; } @@ -6551,6 +6584,10 @@ static int slhdsakey_fors_node_x4(SlhDsaKey* key, const byte* sk_seed, word32 i, adrs, node, state); } #endif + /* Holds FORS secret keys or values derived from them. */ + if (WC_VAR_OK(state)) { + ForceZero(state, sizeof(word64) * SLHDSA_SHAKE_STATE_W); + } WC_FREE_VAR_EX(state, key->heap, DYNAMIC_TYPE_SLHDSA); } @@ -7584,12 +7621,17 @@ int wc_SlhDsaKey_Init(SlhDsaKey* key, enum SlhDsaParam param, void* heap, #ifdef WOLFSSL_SLHDSA_SHA2 if (SLHDSA_IS_SHA2(param)) { - /* Initialize SHA2 hash objects. */ - ret = wc_InitSha256(&key->hash.sha2.sha256); + /* Initialize SHA2 hash objects. The heap hint is passed on, but + * the device id deliberately is not: these objects only ever run + * the inner F, H and PRF compressions, which take the direct block + * path rather than a callback. See SLHDSA_SHA256_RAW_OK(). */ + ret = wc_InitSha256_ex(&key->hash.sha2.sha256, key->heap, + INVALID_DEVID); if (ret == 0) key->hash.sha2.sha256_inited = 1; if ((ret == 0) && (key->params->n > 16)) { - ret = wc_InitSha512(&key->hash.sha2.sha512); + ret = wc_InitSha512_ex(&key->hash.sha2.sha512, key->heap, + INVALID_DEVID); if (ret == 0) key->hash.sha2.sha512_inited = 1; } diff --git a/wolfssl/wolfcrypt/sha256.h b/wolfssl/wolfcrypt/sha256.h index 586ed7b14de..ee7e6201c7b 100644 --- a/wolfssl/wolfcrypt/sha256.h +++ b/wolfssl/wolfcrypt/sha256.h @@ -114,24 +114,32 @@ /* Ports that replace Update and Final with hardware calls do not build * wc_Sha256HashBlock(). The list is negative on purpose: naming a port that - * has it only costs the fast path, missing one is a link error. */ + * has it only costs the fast path, missing one is a link error. A macro a + * port sets only in its own .c reads as undefined here and excludes nothing. + * A FIPS v1 build skips this whole region, leaving the macro undefined and the + * callers on the streaming path, which is the safe direction. */ #if (defined(WOLFSSL_HAVE_LMS) || defined(WOLFSSL_HAVE_SLHDSA)) && \ !defined(WOLFSSL_NO_HASH_RAW) && \ !defined(WOLFSSL_TI_HASH) && \ !defined(WOLFSSL_CRYPTOCELL) && \ - !defined(MAX3266X_SHA) && \ + !defined(WOLFSSL_MAX3266X) && \ + !defined(WOLFSSL_MAX3266X_OLD) && \ !defined(FREESCALE_LTC_SHA) && \ !defined(WOLFSSL_PIC32MZ_HASH) && \ !defined(STM32_HASH_SHA2) && \ - !(defined(WOLFSSL_IMX6_CAAM) && !defined(NO_IMX6_CAAM_HASH)) && \ + !(defined(WOLFSSL_IMX6_CAAM) && !defined(NO_IMX6_CAAM_HASH) && \ + !defined(WOLFSSL_QNX_CAAM)) && \ !(defined(WOLFSSL_SE050) && defined(WOLFSSL_SE050_HASH)) && \ !defined(WOLFSSL_AFALG_HASH) && \ !defined(WOLFSSL_DEVCRYPTO_HASH) && \ - !defined(WOLFSSL_USE_ESP32_CRYPT_HASH_HW) && \ - !defined(WOLFSSL_RENESAS_TSIP_TLS) && \ - !defined(WOLFSSL_RENESAS_TSIP_CRYPTONLY) && \ - !defined(WOLFSSL_RENESAS_SCEPROTECT) && \ - !defined(WOLFSSL_RENESAS_RSIP) && \ + !(defined(WOLFSSL_ESP32_CRYPT) && \ + !defined(NO_WOLFSSL_ESP32_CRYPT_HASH)) && \ + !((defined(WOLFSSL_RENESAS_TSIP_TLS) || \ + defined(WOLFSSL_RENESAS_TSIP_CRYPTONLY)) && \ + !defined(NO_WOLFSSL_RENESAS_TSIP_CRYPT_HASH)) && \ + !((defined(WOLFSSL_RENESAS_SCEPROTECT) || \ + defined(WOLFSSL_RENESAS_RSIP)) && \ + !defined(NO_WOLFSSL_RENESAS_FSPSM_HASH)) && \ !defined(WOLFSSL_RENESAS_RX64_HASH) && \ !defined(PSOC6_HASH_SHA2) && \ !defined(WOLFSSL_IMXRT_DCP) && \ From a0c888bf421d266b58a1d00fde409e1475b8f933 Mon Sep 17 00:00:00 2001 From: Tobias Frauenschlaeger Date: Wed, 19 Aug 2026 13:36:45 +0000 Subject: [PATCH 6/6] Address second review pass on the SLH-DSA hashing work Only the generic wc_Sha256Copy() frees its destination. The KCAPI and AF_ALG ports overwrite it wholesale and allocate a fresh handle, and the generic one returns before its free when a crypto callback takes the copy. Dropping the wc_Sha256Free()/wc_Sha512Free() ahead of the copy therefore stranded a hash handle on every inner hash on those builds. Restore it in all four places and say in the comment why it is needed. wc_InitSha256() takes the crypto callback default device id, not INVALID_DEVID, so initialising the consumer objects with INVALID_DEVID while the midstate objects kept the default made SLHDSA_SHA256_RAW_OK() always true. The direct block path then restored a midstate the callback had never written. Initialise both the same way again and have the guard check the midstate as well, since that is the object whose state words are read. The memory zeroization check registered the block buffer after the hash had run and only when the caller asked for a wipe, so it could not catch an early return and did not match the now unconditional clear. Register it while the data is live and drop the flag. Wipe the WOTS+ chain values at the size of the buffer rather than a parameter-set expression that no longer relates to it, and the FORS nodes at the size of theirs. The obvious length for the nodes, m * n, is wrong there: the Merkle loop halves m once per level, so on the success path it names the width of the level it stopped on rather than the number of entries written, and all but the last few leaf values would survive the free. Also: assert the lane and rate bounds the eight-way helpers assume, taking the rate bound from the widest absorb, H over a tree index, rather than the narrower WOTS+ one, assert the SHA-256 layout the direct path depends on, zero the shared buffer struct at both owners, rename a test that is no longer LMS specific, drop a configure flag that is an alias of one already passed, follow the generated assembly's condition naming, and record the verification stack growth in the ChangeLog. --- .github/configs/pq-all.json | 5 +- ChangeLog.md | 5 +- wolfcrypt/src/sha3_asm.asm | 8 +-- wolfcrypt/src/wc_slhdsa.c | 119 +++++++++++++++++++++++------------- wolfcrypt/test/test.c | 4 +- 5 files changed, 87 insertions(+), 54 deletions(-) diff --git a/.github/configs/pq-all.json b/.github/configs/pq-all.json index de4a8915773..9d0b2725dad 100644 --- a/.github/configs/pq-all.json +++ b/.github/configs/pq-all.json @@ -155,10 +155,9 @@ "--enable-mlkem=make,enc,dec,768", "--disable-qt", "CPPFLAGS=-pedantic -Wdeclaration-after-statement -Wnull-dereference -DWOLFCRYPT_TEST_LINT -DNO_WOLFSSL_CIPHER_SUITE_TEST -DTEST_LIBWOLFSSL_SOURCES_INCLUSION_SEQUENCE"]}, {"name": "slhdsa-only-asm", "minutes": 3, - "comment": "SLH-DSA with the Intel assembly but without ML-KEM/ML-DSA; the 8-way AVX512 Keccak permutation is emitted for those, so this catches SLH-DSA calling a symbol the SHA-3 assembly did not build", + "comment": "SLH-DSA with the Intel assembly but without ML-KEM/ML-DSA; the 8-way AVX512 Keccak permutation is emitted for those, so this catches SLH-DSA calling a symbol the SHA-3 assembly did not build. Linking is checked on any runner, but the 8-way code only executes on an AVX512F+AVX512BW one", "configure": ["--enable-intelasm", "--enable-sp-asm", - "--enable-slhdsa=yes,sha2", "--disable-mlkem", "--disable-mldsa", - "--disable-dilithium"]}, + "--enable-slhdsa=yes,sha2", "--disable-mlkem", "--disable-mldsa"]}, {"name": "slhdsa-no-shake-x8", "minutes": 3, "comment": "SLH-DSA with the Intel assembly but the eight-way AVX512 Keccak path pinned off, so the four-way path is covered whatever the runner CPU; paired with slhdsa-only-asm, a KAT that passes here and fails there isolates the eight-way code", "configure": ["--enable-intelasm", "--enable-sp-asm", "--enable-slhdsa", diff --git a/ChangeLog.md b/ChangeLog.md index ab31caa70f7..1beb1ad1776 100644 --- a/ChangeLog.md +++ b/ChangeLog.md @@ -633,7 +633,10 @@ PR stands for Pull Request, and PR references a GitHub pull request num reuse, instead of allocating one per group of hashes. A SHAKE-128s signature now makes about twelve times fewer allocations. by @Frauschi * Use the 8-way AVX512 Keccak permutation when completing the WOTS+ chains of - an SLH-DSA signature, which speeds up verification. by @Frauschi + an SLH-DSA signature, which speeds up verification. The 8-way verify path + holds its Keccak state and chain values on the stack for the length of one + WOTS+ public key, so peak stack during verification grows by about 2.3 kB on + AVX512 builds; WOLFSSL_SMALL_STACK moves them to the heap. by @Frauschi * Give a whole SLH-DSA XMSS subtree one set of WOTS+ buffers to reuse. A SHAKE-128s signature now makes about 1000 allocations where it made about 144000. by @Frauschi diff --git a/wolfcrypt/src/sha3_asm.asm b/wolfcrypt/src/sha3_asm.asm index fe8178c2150..a55f7415e71 100644 --- a/wolfcrypt/src/sha3_asm.asm +++ b/wolfcrypt/src/sha3_asm.asm @@ -42849,14 +42849,14 @@ sha3_blocksx8_avx512 PROC sha3_blocksx8_avx512 ENDP _TEXT ENDS ENDIF -wc_masm_cond_3b = 0 +wc_masm_cond_4 = 0 IFDEF WOLFSSL_HAVE_MLKEM -wc_masm_cond_3b = 1 +wc_masm_cond_4 = 1 ENDIF IFDEF WOLFSSL_HAVE_MLDSA -wc_masm_cond_3b = 1 +wc_masm_cond_4 = 1 ENDIF -IF wc_masm_cond_3b +IF wc_masm_cond_4 _DATA SEGMENT ALIGN 16 L_sha3_128_blocksx8_seed_avx512_end_mark QWORD 8000000000000000h, 8000000000000000h diff --git a/wolfcrypt/src/wc_slhdsa.c b/wolfcrypt/src/wc_slhdsa.c index edabb8462df..0e1065fa2a1 100644 --- a/wolfcrypt/src/wc_slhdsa.c +++ b/wolfcrypt/src/wc_slhdsa.c @@ -291,6 +291,12 @@ typedef struct SlhDsaWotsBufs { /* Size of an encoded HashAddress. */ #define SLHDSA_HA_SZ 32 +/* The widest 8-way absorb is H over a tree index - PK.seed, the address and a + * 2n-byte message - and it has to fit one SHAKE-256 rate block. */ +wc_static_assert((SLHDSA_MAX_N % 8) == 0); +wc_static_assert((SLHDSA_HA_SZ % 8) == 0); +wc_static_assert(3 * SLHDSA_MAX_N + SLHDSA_HA_SZ < WC_SHA3_256_BLOCK_SIZE); + /* Initialize a HashAddress. * * @param [in] a HashAddress to initialize. @@ -779,14 +785,18 @@ static int slhdsakey_hash_shake_4(wc_Shake* shake, const byte* data1, #endif #ifdef SLHDSA_SHA2_BLOCK_HASH -/* A registered callback expects to see every update, so the direct path is - * only taken for an object no callback has claimed. The key's own hash objects - * are always initialised with INVALID_DEVID, so today this is constant true - * and the direct path always runs; the check is what stops a later change to - * that initialisation from silently bypassing a registered callback. */ +/* The direct path restores the midstate through wc_Sha256's own digest and + * scrubs its block buffer, so a port that changes either layout has to fail + * here rather than write out of bounds. */ +wc_static_assert(sizeof(((wc_Sha256*)0)->buffer) == WC_SHA256_BLOCK_SIZE); +wc_static_assert(sizeof(((wc_Sha256*)0)->digest) == WC_SHA256_DIGEST_SIZE); + +/* A registered callback expects to see every update, and holds the state + * itself, so a claimed midstate has no digest to restore. */ #ifdef WOLF_CRYPTO_CB #define SLHDSA_SHA256_RAW_OK(key) \ - ((key)->hash.sha2.sha256.devId == INVALID_DEVID) + (((key)->hash.sha2.sha256.devId == INVALID_DEVID) && \ + ((key)->hash.sha2.sha256_mid.devId == INVALID_DEVID)) #else #define SLHDSA_SHA256_RAW_OK(key) 1 #endif @@ -888,8 +898,6 @@ wc_static_assert(SLHDSA_HAC_SZ + 32 + 1 + 8 <= WC_SHA256_BLOCK_SIZE); * @param [in] m2_len Length of second message part. * @param [out] hash Buffer to hold hash output. * @param [in] hash_len Number of bytes of hash to output. - * @param [in] zeroize Track the working buffers as secret for the memory - * zeroization check. They are wiped either way. * @return 0 on success. * @return BUFFER_E when the message does not fit the block with its padding. */ @@ -922,6 +930,12 @@ static int slhdsakey_sha256_block_hash(SlhDsaKey* key, const byte* address, c32toa(0, block + WC_SHA256_BLOCK_SIZE - 8); c32toa(bits, block + WC_SHA256_BLOCK_SIZE - 4); + /* Registered while the data is still live, so an early return is caught. */ +#ifdef WOLFSSL_CHECK_MEM_ZERO + wc_MemZero_Add("slhdsa sha256 block", block, WC_SHA256_BLOCK_SIZE); + wc_MemZero_Add("slhdsa sha256 digest", digest, sizeof(digest)); +#endif + /* Restore the midstate and compress. */ XMEMCPY(key->hash.sha2.sha256.digest, key->hash.sha2.sha256_mid.digest, WC_SHA256_DIGEST_SIZE); @@ -930,10 +944,6 @@ static int slhdsakey_sha256_block_hash(SlhDsaKey* key, const byte* address, XMEMCPY(hash, digest, hash_len); } -#ifdef WOLFSSL_CHECK_MEM_ZERO - wc_MemZero_Add("slhdsa sha256 block", block, WC_SHA256_BLOCK_SIZE); - wc_MemZero_Add("slhdsa sha256 digest", digest, sizeof(digest)); -#endif /* Cleared unconditionally rather than per call site. */ ForceZero(block, WC_SHA256_BLOCK_SIZE); ForceZero(digest, sizeof(digest)); @@ -965,7 +975,12 @@ static int slhdsakey_sha256_api_hash(SlhDsaKey* key, const byte* address, int ret; byte digest[WC_SHA256_DIGEST_SIZE]; - /* Restore the midstate. wc_Sha256Copy() releases the destination. */ + /* Only the generic wc_Sha256Copy() frees its destination; the KCAPI, AF_ALG + * and crypto callback copies would strand a handle per hash. */ + if (key->hash.sha2.sha256_inited) { + wc_Sha256Free(&key->hash.sha2.sha256); + key->hash.sha2.sha256_inited = 0; + } ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, &key->hash.sha2.sha256); if (ret == 0) { key->hash.sha2.sha256_inited = 1; @@ -1040,7 +1055,11 @@ static int slhdsakey_sha512_hash(SlhDsaKey* key, const byte* address, int ret; byte digest[WC_SHA512_DIGEST_SIZE]; - /* Restore the midstate. wc_Sha512Copy() releases the destination. */ + /* Release the previous state first - see slhdsakey_sha256_api_hash(). */ + if (key->hash.sha2.sha512_inited) { + wc_Sha512Free(&key->hash.sha2.sha512); + key->hash.sha2.sha512_inited = 0; + } ret = wc_Sha512Copy(&key->hash.sha2.sha512_mid, &key->hash.sha2.sha512); if (ret == 0) { key->hash.sha2.sha512_inited = 1; @@ -1227,6 +1246,10 @@ static int slhdsakey_hash_start_addr_sha2(SlhDsaKey* key, if (n == WC_SLHDSA_N_128) { /* Category 1: SHA-256 -- use sha256_2 (T_l must not collide with * sha256 which is used by F and H). */ + if (key->hash.sha2.sha256_2_inited) { + wc_Sha256Free(&key->hash.sha2.sha256_2); + key->hash.sha2.sha256_2_inited = 0; + } ret = wc_Sha256Copy(&key->hash.sha2.sha256_mid, &key->hash.sha2.sha256_2); if (ret == 0) { @@ -1238,6 +1261,10 @@ static int slhdsakey_hash_start_addr_sha2(SlhDsaKey* key, else { /* Categories 3, 5: SHA-512 -- use sha512_2 (T_l must not collide * with sha512 which is used by H). */ + if (key->hash.sha2.sha512_2_inited) { + wc_Sha512Free(&key->hash.sha2.sha512_2); + key->hash.sha2.sha512_2_inited = 0; + } ret = wc_Sha512Copy(&key->hash.sha2.sha512_mid, &key->hash.sha2.sha512_2); if (ret == 0) { @@ -2725,6 +2752,7 @@ static void slhdsakey_shake256_set_chain_addr_idx_x8(word64* state, word32 o, * @param [in] fixed Caller owned state head, already filled. * @param [in] state Caller owned x8 Keccak state. * @return 0 on success. + * @return Error code from saving the vector registers. */ static int slhdsakey_chain_idx_x8(byte* sk, word32 i, word32 s, byte n, word32 o, word64* fixed, word64* state) @@ -2783,7 +2811,7 @@ static int slhdsakey_chain_idx_x8(byte* sk, word32 i, word32 s, byte n, * @param [in] n Number of bytes in hash output. * @param [in] ca Chain address start index. * @param [out] sk Buffer to hold hash output. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. * @return SHAKE-256 error return code on digest failure. */ @@ -2825,8 +2853,8 @@ static int slhdsakey_hash_prf_x4(const byte* pk_seed, const byte* sk_seed, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. - * @param [in] fixed Caller owned buffer for the unchanging state head. - * @param [in] state Caller owned Keccak state. + * @param [in, out] fixed Caller owned buffer for the unchanging state head. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_chain_x4_16(byte* sk, const byte* pk_seed, byte* addr, @@ -2879,8 +2907,8 @@ static int slhdsakey_chain_x4_16(byte* sk, const byte* pk_seed, byte* addr, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. - * @param [in] fixed Caller owned buffer for the unchanging state head. - * @param [in] state Caller owned Keccak state. + * @param [in, out] fixed Caller owned buffer for the unchanging state head. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_chain_x4_24(byte* sk, const byte* pk_seed, byte* addr, @@ -2933,8 +2961,8 @@ static int slhdsakey_chain_x4_24(byte* sk, const byte* pk_seed, byte* addr, * @param [in] pk_seed Public key seed. * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. - * @param [in] fixed Caller owned buffer for the unchanging state head. - * @param [in] state Caller owned Keccak state. + * @param [in, out] fixed Caller owned buffer for the unchanging state head. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_chain_x4_32(byte* sk, const byte* pk_seed, byte* addr, @@ -3337,7 +3365,7 @@ static int slhdsakey_wots_pkgen_chain_x4_16(SlhDsaKey* key, const byte* sk_seed, * with public chain values. The x4 PRF fills up to a 4-lane multiple * (beyond len), so wipe the whole buffer. */ if (ret != 0) { - ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 16); + ForceZero(sk, SLHDSA_WOTS_SK_SZ); } return ret; } @@ -3410,7 +3438,7 @@ static int slhdsakey_wots_pkgen_chain_x4_24(SlhDsaKey* key, const byte* sk_seed, * with public chain values. The x4 PRF fills up to a 4-lane multiple * (beyond len), so wipe the whole buffer. */ if (ret != 0) { - ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 24); + ForceZero(sk, SLHDSA_WOTS_SK_SZ); } return ret; } @@ -3482,7 +3510,7 @@ static int slhdsakey_wots_pkgen_chain_x4_32(SlhDsaKey* key, const byte* sk_seed, /* On error sk still holds secret WOTS+ leaves, and the x4 PRF fills past * len to a 4-lane multiple, so wipe the whole buffer. */ if (ret != 0) { - ForceZero(sk, (SLHDSA_MAX_MSG_SZ + 3) * 32); + ForceZero(sk, SLHDSA_WOTS_SK_SZ); } return ret; } @@ -3604,9 +3632,10 @@ static int slhdsakey_hash_prf_x8(const byte* pk_seed, const byte* sk_seed, * @param [in] addr Encoded HashAddress. * @param [in] ca Chain address start index. * @param [in] n Number of bytes in each hash. - * @param [in] fixed Caller owned buffer for the unchanging state head. + * @param [in, out] fixed Caller owned buffer for the unchanging state head. * @param [in] state Caller owned x8 Keccak state. * @return 0 on success. + * @return Error code from saving the vector registers. */ static int slhdsakey_chain_x8(byte* sk, const byte* pk_seed, byte* addr, byte ca, byte n, word64* fixed, word64* state) @@ -5193,8 +5222,9 @@ static int slhdsakey_wots_pk_from_sig(SlhDsaKey* key, const byte* sig, * @param [in, out] adrs HashAddress - WOTS HASH. * @param [out] node Root node. * @param [in] bufs Buffers shared by the WOTS+ public keys of one - * subtree. + * subtree. bufs->sk holds SLHDSA_WOTS_SK_SZ bytes. * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, @@ -5305,8 +5335,9 @@ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, * @param [in, out] adrs HashAddress - WOTS HASH. * @param [out] node Root node. * @param [in] bufs Buffers shared by the WOTS+ public keys of one - * subtree. + * subtree. bufs->sk holds SLHDSA_WOTS_SK_SZ bytes. * @return 0 on success. + * @return MEMORY_E on dynamic memory allocation failure. * @return SHAKE-256 error return code on digest failure. */ static int slhdsakey_xmss_node(SlhDsaKey* key, const byte* sk_seed, int i, @@ -5376,6 +5407,7 @@ static int slhdsakey_root_from_seed(SlhDsaKey* key, byte* root) WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_STATE_W, key->heap); #endif + XMEMSET(&bufs, 0, sizeof(bufs)); #ifdef SLHDSA_NEED_WOTS_SK_BUF WC_ALLOC_VAR_EX(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); @@ -5461,6 +5493,7 @@ static int slhdsakey_xmss_sign(SlhDsaKey* key, const byte* m, WC_DECLARE_VAR(state, word64, SLHDSA_SHAKE_STATE_W, key->heap); #endif + XMEMSET(&bufs, 0, sizeof(bufs)); #ifdef SLHDSA_NEED_WOTS_SK_BUF WC_ALLOC_VAR_EX(sk, byte, SLHDSA_WOTS_SK_SZ, key->heap, DYNAMIC_TYPE_SLHDSA, ret = MEMORY_E); @@ -5859,7 +5892,7 @@ static int slhdsakey_fors_sk_gen(SlhDsaKey* key, const byte* sk_seed, * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. * @param [out] node Buffer to hold hash output. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_hash_prf_ti_x4(const byte* pk_seed, const byte* sk_seed, @@ -5900,7 +5933,7 @@ static int slhdsakey_hash_prf_ti_x4(const byte* pk_seed, const byte* sk_seed, * @param [in, out] node On in, n-byte messages. On out, n-byte outputs. * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, @@ -5950,7 +5983,7 @@ static int slhdsakey_hash_f_ti_x4(const byte* pk_seed, byte* addr, byte* node, * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. * @param [out] hash Buffer to hold hash output. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_hash_h_ti_x4(const byte* pk_seed, byte* addr, @@ -5992,7 +6025,7 @@ static int slhdsakey_hash_h_ti_x4(const byte* pk_seed, byte* addr, * @param [in] n Number of bytes in each hash. * @param [in] ti Tree index start value. * @param [out] node Eight n-byte outputs. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_hash_prf_ti_x8(const byte* pk_seed, const byte* sk_seed, @@ -6024,7 +6057,7 @@ static int slhdsakey_hash_prf_ti_x8(const byte* pk_seed, const byte* sk_seed, * @param [in, out] node On in, eight n-byte messages. On out, the outputs. * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_hash_f_ti_x8(const byte* pk_seed, byte* addr, byte* node, @@ -6065,7 +6098,7 @@ static int slhdsakey_hash_f_ti_x8(const byte* pk_seed, byte* addr, byte* node, * @param [in] n Number of bytes in hash output. * @param [in] ti Tree index start value. * @param [out] hash Buffer to hold eight n-byte hash outputs. - * @param [in] state Caller owned Keccak state. + * @param [in, out] state Caller owned Keccak state. * @return 0 on success. */ static int slhdsakey_hash_h_ti_x8(const byte* pk_seed, byte* addr, @@ -6386,7 +6419,8 @@ static int slhdsakey_fors_node_x4_low(SlhDsaKey* key, const byte* sk_seed, ret = HASH_H(key, pk_seed, adrs, nodes, n, node); } - /* Holds FORS secret keys or values derived from them. */ + /* Holds FORS secret keys or values derived from them. Wiped by buffer + * size, as the Merkle loop leaves m at the last level's width. */ if (WC_VAR_OK(nodes)) { ForceZero(nodes, (1 << SLHDSA_MAX_FORS_NODE_DEPTH) * SLHDSA_MAX_N); } @@ -6435,7 +6469,7 @@ static int slhdsakey_fors_node_x4_high(SlhDsaKey* key, const byte* sk_seed, byte n = key->params->n; word32 j; word32 z2 = z % SLHDSA_MAX_FORS_NODE_DEPTH; - word32 m; + word32 m = 0; WC_DECLARE_VAR(nodes, byte, (1 << SLHDSA_MAX_FORS_NODE_TOP_DEPTH) * SLHDSA_MAX_N, key->heap); @@ -6514,7 +6548,8 @@ static int slhdsakey_fors_node_x4_high(SlhDsaKey* key, const byte* sk_seed, ret = HASH_H(key, pk_seed, adrs, nodes, n, node); } - /* Holds FORS secret keys or values derived from them. */ + /* Holds FORS secret keys or values derived from them. Wiped by buffer + * size, as the Merkle loop leaves m at the last level's width. */ if (WC_VAR_OK(nodes)) { ForceZero(nodes, (1 << SLHDSA_MAX_FORS_NODE_TOP_DEPTH) * SLHDSA_MAX_N); } @@ -7621,17 +7656,13 @@ int wc_SlhDsaKey_Init(SlhDsaKey* key, enum SlhDsaParam param, void* heap, #ifdef WOLFSSL_SLHDSA_SHA2 if (SLHDSA_IS_SHA2(param)) { - /* Initialize SHA2 hash objects. The heap hint is passed on, but - * the device id deliberately is not: these objects only ever run - * the inner F, H and PRF compressions, which take the direct block - * path rather than a callback. See SLHDSA_SHA256_RAW_OK(). */ - ret = wc_InitSha256_ex(&key->hash.sha2.sha256, key->heap, - INVALID_DEVID); + /* Same device id as the midstate objects they are copied from, + * so SLHDSA_SHA256_RAW_OK() sees a consistent pair. */ + ret = wc_InitSha256(&key->hash.sha2.sha256); if (ret == 0) key->hash.sha2.sha256_inited = 1; if ((ret == 0) && (key->params->n > 16)) { - ret = wc_InitSha512_ex(&key->hash.sha2.sha512, key->heap, - INVALID_DEVID); + ret = wc_InitSha512(&key->hash.sha2.sha512); if (ret == 0) key->hash.sha2.sha512_inited = 1; } diff --git a/wolfcrypt/test/test.c b/wolfcrypt/test/test.c index 8e32a244ca2..ce8a108cc40 100644 --- a/wolfcrypt/test/test.c +++ b/wolfcrypt/test/test.c @@ -6522,7 +6522,7 @@ static wc_test_ret_t sha256_large_hash_test(wc_Sha256* sha) #endif /* NO_LARGE_HASH_TEST */ #ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK -static wc_test_ret_t sha256_lms_test(wc_Sha256* sha) +static wc_test_ret_t sha256_hash_block_test(wc_Sha256* sha) { byte hash[WC_SHA256_DIGEST_SIZE]; wc_test_ret_t ret = 0; @@ -6615,7 +6615,7 @@ WOLFSSL_TEST_SUBROUTINE wc_test_ret_t sha256_test(void) return ret; #endif #ifdef WOLFSSL_HAVE_SHA256_HASH_BLOCK - if ((ret = sha256_lms_test(&sha)) != 0) + if ((ret = sha256_hash_block_test(&sha)) != 0) return ret; #endif #if !defined(HAVE_SELFTEST) && (!defined(HAVE_FIPS) || FIPS_VERSION_GE(7, 0))