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)) ||