Split MLKEM's scalar_*code functions into compile-time alternatives by bit size. Gains 10% on AMD EPYC 7B13 and 16% on Apple M1 Pro. Bug: 503700354 Change-Id: Icd8fc6e9b0eb270b1fa21e159521ba326a6a6964 Reviewed-on: https://boringssl-review.googlesource.com/c/boringssl/+/92828 Reviewed-by: Xiangfei Ding <xfding@google.com> Commit-Queue: Rudolf Polzer <rpolzer@google.com> Presubmit-BoringSSL-Verified: boringssl-scoped@luci-project-accounts.iam.gserviceaccount.com <boringssl-scoped@luci-project-accounts.iam.gserviceaccount.com>
diff --git a/crypto/fipsmodule/mlkem/mlkem.cc.inc b/crypto/fipsmodule/mlkem/mlkem.cc.inc index 107b589..7c86c6b 100644 --- a/crypto/fipsmodule/mlkem/mlkem.cc.inc +++ b/crypto/fipsmodule/mlkem/mlkem.cc.inc
@@ -477,126 +477,267 @@ } } -const uint8_t kMasks[8] = {0x01, 0x03, 0x07, 0x0f, 0x1f, 0x3f, 0x7f, 0xff}; +// Encodes a scalar of 256 |BITS|-bit words into 32*|BITS| bytes by splitting +// and joining into bytes using LSB-first bit order (i.e. opposite to standard +// reading order). See below for examples. If an input is >= 1 << |BITS|, the +// result is undefined. +template <int BITS> +void scalar_encode(uint8_t *out, const scalar *s); -void scalar_encode(uint8_t *out, const scalar *s, int bits) { - assert(bits <= (int)sizeof(*s->c) * 8 && bits != 1); - - uint8_t out_byte = 0; - int out_byte_bits = 0; - - for (int i = 0; i < DEGREE; i++) { - uint16_t element = s->c[i]; - int element_bits_done = 0; - - while (element_bits_done < bits) { - int chunk_bits = bits - element_bits_done; - int out_bits_remaining = 8 - out_byte_bits; - if (chunk_bits >= out_bits_remaining) { - chunk_bits = out_bits_remaining; - out_byte |= (element & kMasks[chunk_bits - 1]) << out_byte_bits; - *out = out_byte; - out++; - out_byte_bits = 0; - out_byte = 0; - } else { - out_byte |= (element & kMasks[chunk_bits - 1]) << out_byte_bits; - out_byte_bits += chunk_bits; - } - - element_bits_done += chunk_bits; - element >>= chunk_bits; - } - } - - if (out_byte_bits > 0) { - *out = out_byte; +// Encodes a scalar of 256 10-bit words into 320 bytes as follows: +// 000000Aaaaaaaaaa 000000Bbbbbbbbbb 000000Cccccccccc 000000Dddddddddd ... +// -> aaaaaaaa bbbbbbAa ccccBbbb ddCccccc Dddddddd ... +template <> +void scalar_encode<10>(uint8_t out[320], const scalar *s) { + for (int i = 0; i < DEGREE; i += 4) { + uint16_t s0 = s->c[i]; + uint16_t s1 = s->c[i + 1]; + uint16_t s2 = s->c[i + 2]; + uint16_t s3 = s->c[i + 3]; + declassify_assert((s0 | s1 | s2 | s3) < (1 << 10)); + out[0] = (uint8_t)s0; + out[1] = (uint8_t)((s0 >> 8) | (s1 << 2)); + out[2] = (uint8_t)((s1 >> 6) | (s2 << 4)); + out[3] = (uint8_t)((s2 >> 4) | (s3 << 6)); + out[4] = (uint8_t)(s3 >> 2); + out += 5; } } -// scalar_encode_1 is |scalar_encode| specialised for |bits| == 1. -void scalar_encode_1(uint8_t out[32], const scalar *s) { +// Encodes a scalar of 256 12-bit words into 384 bytes as follows: +// 0000Aaaaaaaaaaaa 0000Bbbbbbbbbbbb 0000Cccccccccccc 0000Dddddddddddd ... +// -> aaaaaaaa bbbbAaaa Bbbbbbbb cccccccc ddddCccc Dddddddd .... +template <> +void scalar_encode<12>(uint8_t out[384], const scalar *s) { + for (int i = 0; i < DEGREE; i += 2) { + uint16_t s0 = s->c[i]; + uint16_t s1 = s->c[i + 1]; + declassify_assert((s0 | s1) < (1 << 12)); + out[0] = (uint8_t)s0; + out[1] = (uint8_t)((s0 >> 8) | (s1 << 4)); + out[2] = (uint8_t)(s1 >> 4); + out += 3; + } +} + +// Encodes a scalar of 256 4-bit words into 128 bytes as follows: +// 000000000000Aaaa 00000000000Bbbb 000000000000Cccc 000000000000Dddd ... +// -> BbbbAaaa DdddCccc ... +template <> +void scalar_encode<4>(uint8_t out[128], const scalar *s) { + for (int i = 0; i < DEGREE; i += 2) { + uint16_t s0 = s->c[i]; + uint16_t s1 = s->c[i + 1]; + declassify_assert((s0 | s1) < (1 << 4)); + out[0] = (uint8_t)(s0 | (s1 << 4)); + out += 1; + } +} + +// Encodes a scalar of 256 11-bit words into 352 bytes as follows: +// 00000Aaaaaaaaaaa 00000Bbbbbbbbbbb 00000Ccccccccccc 00000Ddddddddddd ... +// -> aaaaaaaa bbbbbAaa ccBbbbbb cccccccc dddddddC eeeeDddd fEeeeeee ... +template <> +void scalar_encode<11>(uint8_t out[352], const scalar *s) { + for (int i = 0; i < DEGREE; i += 8) { + uint16_t s0 = s->c[i]; + uint16_t s1 = s->c[i + 1]; + uint16_t s2 = s->c[i + 2]; + uint16_t s3 = s->c[i + 3]; + uint16_t s4 = s->c[i + 4]; + uint16_t s5 = s->c[i + 5]; + uint16_t s6 = s->c[i + 6]; + uint16_t s7 = s->c[i + 7]; + declassify_assert((s0 | s1 | s2 | s3 | s4 | s5 | s6 | s7) < (1 << 11)); + out[0] = (uint8_t)s0; + out[1] = (uint8_t)((s0 >> 8) | (s1 << 3)); + out[2] = (uint8_t)((s1 >> 5) | (s2 << 6)); + out[3] = (uint8_t)(s2 >> 2); + out[4] = (uint8_t)((s2 >> 10) | (s3 << 1)); + out[5] = (uint8_t)((s3 >> 7) | (s4 << 4)); + out[6] = (uint8_t)((s4 >> 4) | (s5 << 7)); + out[7] = (uint8_t)(s5 >> 1); + out[8] = (uint8_t)((s5 >> 9) | (s6 << 2)); + out[9] = (uint8_t)((s6 >> 6) | (s7 << 5)); + out[10] = (uint8_t)(s7 >> 3); + out += 11; + } +} + +// Encodes a scalar of 256 5-bit words into 160 bytes as follows: +// 00000000000Aaaaa 00000000000Bbbbb 00000000000Ccccc 00000000000Ddddd ... +// -> bbbAaaaa dCccccBb eeeeDddd ggFffffE HhhhhGgg ... +template <> +void scalar_encode<5>(uint8_t out[160], const scalar *s) { + for (int i = 0; i < DEGREE; i += 8) { + uint16_t s0 = s->c[i]; + uint16_t s1 = s->c[i + 1]; + uint16_t s2 = s->c[i + 2]; + uint16_t s3 = s->c[i + 3]; + uint16_t s4 = s->c[i + 4]; + uint16_t s5 = s->c[i + 5]; + uint16_t s6 = s->c[i + 6]; + uint16_t s7 = s->c[i + 7]; + declassify_assert((s0 | s1 | s2 | s3 | s4 | s5 | s6 | s7) < (1 << 5)); + out[0] = (uint8_t)(s0 | (s1 << 5)); + out[1] = (uint8_t)((s1 >> 3) | (s2 << 2) | (s3 << 7)); + out[2] = (uint8_t)((s3 >> 1) | (s4 << 4)); + out[3] = (uint8_t)((s4 >> 4) | (s5 << 1) | (s6 << 6)); + out[4] = (uint8_t)((s6 >> 2) | (s7 << 3)); + out += 5; + } +} + +// Encodes a scalar of 256 1-bit "words" into 32 bytes as follows: +// 0000000000000A 000000000000000B 000000000000000C 000000000000000D ... +// -> HGFEDCBA PONMLKJI XWVUTSRQ ... +// This order is best understood as the natural way of joining into bytes +// assuming LSB-first bit order. +template <> +void scalar_encode<1>(uint8_t out[32], const scalar *s) { for (int i = 0; i < DEGREE; i += 8) { uint8_t out_byte = 0; for (int j = 0; j < 8; j++) { - out_byte |= (s->c[i + j] & 1) << j; + declassify_assert(s->c[i + j] <= 1); + out_byte |= s->c[i + j] << j; } - *out = out_byte; - out++; + out[i / 8] = out_byte; } } // Encodes an entire vector into 32*|RANK|*|bits| bytes. Note that since 256 // (DEGREE) is divisible by 8, the individual vector entries will always fill a // whole number of bytes, so we do not need to worry about bit packing here. -template <int RANK> -void vector_encode(uint8_t *out, const vector<RANK> *a, int bits) { +template <int bits, int RANK> +void vector_encode(uint8_t *out, const vector<RANK> *a) { for (int i = 0; i < RANK; i++) { - scalar_encode(out + i * bits * DEGREE / 8, &a->v[i], bits); + scalar_encode<bits>(out + i * bits * DEGREE / 8, &a->v[i]); } } +// The inverse of |scalar_encode|. Returns 1 iff the encoded scalar is valid, +// i.e. all components are < |kPrime|. Otherwise, returns 0 and the value of +// |out| is undefined. +template <int BITS> +int scalar_decode(scalar *out, const uint8_t *in); + +template <> +int scalar_decode<10>(scalar *out, const uint8_t in[320]) { + for (int i = 0; i < DEGREE; i += 4) { + uint16_t s0 = (uint16_t)(in[0] | ((in[1] & 0x03) << 8)); + uint16_t s1 = (uint16_t)((in[1] >> 2) | ((in[2] & 0x0f) << 6)); + uint16_t s2 = (uint16_t)((in[2] >> 4) | ((in[3] & 0x3f) << 4)); + uint16_t s3 = (uint16_t)((in[3] >> 6) | (in[4] << 2)); + out->c[i] = s0; + out->c[i + 1] = s1; + out->c[i + 2] = s2; + out->c[i + 3] = s3; + in += 5; + } + return 1; +} + +template <> +int scalar_decode<12>(scalar *out, const uint8_t in[384]) { + for (int i = 0; i < DEGREE; i += 2) { + uint16_t s0 = (uint16_t)(in[0] | ((in[1] & 0x0f) << 8)); + uint16_t s1 = (uint16_t)((in[1] >> 4) | (in[2] << 4)); + if (constant_time_declassify_int((s0 | s1) >= kPrime)) { + if (s0 >= kPrime || s1 >= kPrime) { + return 0; + } + } + out->c[i] = s0; + out->c[i + 1] = s1; + in += 3; + } + return 1; +} + +template <> +int scalar_decode<4>(scalar *out, const uint8_t in[128]) { + for (int i = 0; i < DEGREE; i += 2) { + uint16_t s0 = (uint16_t)(in[0] & 0x0f); + uint16_t s1 = (uint16_t)(in[0] >> 4); + // kPrime is 3329, so 4-bit values are always < kPrime. + out->c[i] = s0; + out->c[i + 1] = s1; + in += 1; + } + return 1; +} + // scalar_decode parses |DEGREE * bits| bits from |in| into |DEGREE| values in // |out|. It returns one on success and zero if any parsed value is >= // |kPrime|. -int scalar_decode(scalar *out, const uint8_t *in, int bits) { - assert(bits <= (int)sizeof(*out->c) * 8 && bits != 1); - - uint8_t in_byte = 0; - int in_byte_bits_left = 0; - - for (int i = 0; i < DEGREE; i++) { - uint16_t element = 0; - int element_bits_done = 0; - - while (element_bits_done < bits) { - if (in_byte_bits_left == 0) { - in_byte = *in; - in++; - in_byte_bits_left = 8; - } - - int chunk_bits = bits - element_bits_done; - if (chunk_bits > in_byte_bits_left) { - chunk_bits = in_byte_bits_left; - } - - element |= (in_byte & kMasks[chunk_bits - 1]) << element_bits_done; - in_byte_bits_left -= chunk_bits; - in_byte >>= chunk_bits; - - element_bits_done += chunk_bits; - } - - // An element is only out of range in the case of invalid input, in which - // case it is okay to leak the comparison. - if (constant_time_declassify_int(element >= kPrime)) { - return 0; - } - out->c[i] = element; +template <> +int scalar_decode<11>(scalar *out, const uint8_t in[352]) { + for (int i = 0; i < DEGREE; i += 8) { + uint16_t s0 = (uint16_t)(in[0] | ((in[1] & 0x07) << 8)); + uint16_t s1 = (uint16_t)((in[1] >> 3) | ((in[2] & 0x3f) << 5)); + uint16_t s2 = + (uint16_t)((in[2] >> 6) | (in[3] << 2) | ((in[4] & 0x01) << 10)); + uint16_t s3 = (uint16_t)((in[4] >> 1) | ((in[5] & 0x0f) << 7)); + uint16_t s4 = (uint16_t)((in[5] >> 4) | ((in[6] & 0x7f) << 4)); + uint16_t s5 = + (uint16_t)((in[6] >> 7) | (in[7] << 1) | ((in[8] & 0x03) << 9)); + uint16_t s6 = (uint16_t)((in[8] >> 2) | ((in[9] & 0x1f) << 6)); + uint16_t s7 = (uint16_t)((in[9] >> 5) | (in[10] << 3)); + out->c[i] = s0; + out->c[i + 1] = s1; + out->c[i + 2] = s2; + out->c[i + 3] = s3; + out->c[i + 4] = s4; + out->c[i + 5] = s5; + out->c[i + 6] = s6; + out->c[i + 7] = s7; + in += 11; } - return 1; } -// scalar_decode_1 is |scalar_decode| specialised for |bits| == 1. -void scalar_decode_1(scalar *out, const uint8_t in[32]) { +template <> +int scalar_decode<5>(scalar *out, const uint8_t in[160]) { for (int i = 0; i < DEGREE; i += 8) { - uint8_t in_byte = *in; - in++; + uint16_t s0 = (uint16_t)(in[0] & 0x1f); + uint16_t s1 = (uint16_t)((in[0] >> 5) | ((in[1] & 0x03) << 3)); + uint16_t s2 = (uint16_t)((in[1] >> 2) & 0x1f); + uint16_t s3 = (uint16_t)((in[1] >> 7) | ((in[2] & 0x0f) << 1)); + uint16_t s4 = (uint16_t)((in[2] >> 4) | ((in[3] & 0x01) << 4)); + uint16_t s5 = (uint16_t)((in[3] >> 1) & 0x1f); + uint16_t s6 = (uint16_t)((in[3] >> 6) | ((in[4] & 0x07) << 2)); + uint16_t s7 = (uint16_t)(in[4] >> 3); + // kPrime is 3329, so 5-bit values are always < kPrime. + out->c[i] = s0; + out->c[i + 1] = s1; + out->c[i + 2] = s2; + out->c[i + 3] = s3; + out->c[i + 4] = s4; + out->c[i + 5] = s5; + out->c[i + 6] = s6; + out->c[i + 7] = s7; + in += 5; + } + return 1; +} + +template <> +int scalar_decode<1>(scalar *out, const uint8_t in[32]) { + for (int i = 0; i < DEGREE; i += 8) { + uint8_t in_byte = in[i / 8]; for (int j = 0; j < 8; j++) { - out->c[i + j] = in_byte & 1; - in_byte >>= 1; + out->c[i + j] = (in_byte >> j) & 1; } } + return 1; } // Decodes 32*|RANK|*|bits| bytes from |in| into |out|. It returns one on // success or zero if any parsed value is >= |kPrime|. -template <int RANK> -inline int vector_decode(vector<RANK> *out, const uint8_t *in, int bits) { +template <int bits, int RANK> +inline int vector_decode(vector<RANK> *out, const uint8_t *in) { for (int i = 0; i < RANK; i++) { - if (!scalar_decode(&out->v[i], in + i * bits * DEGREE / 8, bits)) { + if (!scalar_decode<bits>(&out->v[i], in + i * bits * DEGREE / 8)) { return 0; } } @@ -690,18 +831,18 @@ constexpr int dv = RANK == RANK768 ? kDV768 : kDV1024; vector<RANK> u; - vector_decode(&u, ciphertext, du); + vector_decode<du>(&u, ciphertext); vector_decompress(&u, du); vector_ntt(&u); scalar v; - scalar_decode(&v, ciphertext + compressed_vector_size(RANK), dv); + scalar_decode<dv>(&v, ciphertext + compressed_vector_size(RANK)); scalar_decompress(&v, dv); scalar mask; scalar_inner_product(&mask, &priv->s, &u); scalar_inverse_ntt(&mask); scalar_sub(&v, &mask); scalar_compress(&v, 1); - scalar_encode_1(out, &v); + scalar_encode<1>(out, &v); } template <int RANK> @@ -711,7 +852,7 @@ if (!CBB_add_space(out, &vector_output, encoded_vector_size(RANK))) { return bcm_status::failure; } - vector_encode(vector_output, &pub->t, kLog2Prime); + vector_encode<kLog2Prime>(vector_output, &pub->t); if (!CBB_add_bytes(out, pub->rho, sizeof(pub->rho))) { return bcm_status::failure; } @@ -801,13 +942,13 @@ scalar_inverse_ntt(&v); scalar_add(&v, &scalar_error); mlkem::scalar expanded_message; - scalar_decode_1(&expanded_message, message); + scalar_decode<1>(&expanded_message, message); scalar_decompress(&expanded_message, 1); scalar_add(&v, &expanded_message); vector_compress(&u, du); - vector_encode(out, &u, du); + vector_encode<du>(out, &u); scalar_compress(&v, dv); - scalar_encode(out + mlkem::compressed_vector_size(RANK), &v, dv); + scalar_encode<dv>(out + mlkem::compressed_vector_size(RANK), &v); } // See section 6.3 @@ -851,7 +992,7 @@ int mlkem_parse_public_key_no_hash(public_key<RANK> *pub, CBS *in) { CBS t_bytes; if (!CBS_get_bytes(in, &t_bytes, encoded_vector_size(RANK)) || - !vector_decode(&pub->t, CBS_data(&t_bytes), kLog2Prime) || + !vector_decode<kLog2Prime>(&pub->t, CBS_data(&t_bytes)) || !CBS_copy_bytes(in, pub->rho, sizeof(pub->rho))) { return 0; } @@ -874,7 +1015,7 @@ int mlkem_parse_private_key(private_key<RANK> *priv, CBS *in) { CBS s_bytes; if (!CBS_get_bytes(in, &s_bytes, encoded_vector_size(RANK)) || - !vector_decode(&priv->s, CBS_data(&s_bytes), kLog2Prime) || + !vector_decode<kLog2Prime>(&priv->s, CBS_data(&s_bytes)) || !mlkem_parse_public_key_no_hash(&priv->pub, in) || !CBS_copy_bytes(in, priv->pub.public_key_hash, sizeof(priv->pub.public_key_hash)) || @@ -892,7 +1033,7 @@ if (!CBB_add_space(out, &s_output, encoded_vector_size(RANK))) { return 0; } - vector_encode(s_output, &priv->s, kLog2Prime); + vector_encode<kLog2Prime>(s_output, &priv->s); if (!bcm_success(mlkem_marshal_public_key(out, &priv->pub)) || !CBB_add_bytes(out, priv->pub.public_key_hash, sizeof(priv->pub.public_key_hash)) ||