Clean up mlkem_test.cc to use a traits object

Change-Id: I705202745648d3ff42325362a54029771a5ff024
Reviewed-on: https://boringssl-review.googlesource.com/c/boringssl/+/98111
Commit-Queue: David Benjamin <davidben@google.com>
Reviewed-by: Adam Langley <agl@google.com>
diff --git a/crypto/mlkem/mlkem_test.cc b/crypto/mlkem/mlkem_test.cc
index 6a94c26..03c220b 100644
--- a/crypto/mlkem/mlkem_test.cc
+++ b/crypto/mlkem/mlkem_test.cc
@@ -34,6 +34,43 @@
 BSSL_NAMESPACE_BEGIN
 namespace {
 
+// Arguments are templated to avoid needing to repeat the underlying function's
+// type signature. We cannot use p_mldsa.cc's pattern because, on Windows,
+// cross-dll function pointers are not constexpr.
+#define TRAIT_METHOD(method_name, function_name) \
+  template <typename... Args>                    \
+  static auto method_name(Args... args) {        \
+    return function_name(args...);               \
+  }
+
+#define MAKE_MLKEM_TRAITS(bits)                                                \
+  struct MLKEM##bits##Traits {                                                 \
+    using PublicKey = MLKEM##bits##_public_key;                                \
+    using PrivateKey = MLKEM##bits##_private_key;                              \
+                                                                               \
+    static constexpr size_t kPublicKeyBytes = MLKEM##bits##_PUBLIC_KEY_BYTES;  \
+    static constexpr size_t kPrivateKeyBytes =                                 \
+        BCM_MLKEM##bits##_PRIVATE_KEY_BYTES;                                   \
+    static constexpr size_t kCiphertextBytes = MLKEM##bits##_CIPHERTEXT_BYTES; \
+                                                                               \
+    TRAIT_METHOD(GenerateKey, MLKEM##bits##_generate_key)                      \
+    TRAIT_METHOD(PrivateKeyFromSeed, MLKEM##bits##_private_key_from_seed)      \
+    TRAIT_METHOD(PublicFromPrivate, MLKEM##bits##_public_from_private)         \
+    TRAIT_METHOD(ParsePublicKey, MLKEM##bits##_parse_public_key)               \
+    TRAIT_METHOD(MarshalPublicKey, MLKEM##bits##_marshal_public_key)           \
+    TRAIT_METHOD(ParsePrivateKey, BCM_mlkem##bits##_parse_private_key)         \
+    TRAIT_METHOD(MarshalPrivateKey, BCM_mlkem##bits##_marshal_private_key)     \
+    TRAIT_METHOD(GenerateKeyExternalSeed,                                      \
+                 BCM_mlkem##bits##_generate_key_external_seed)                 \
+    TRAIT_METHOD(EncapExternalEntropy,                                         \
+                 BCM_mlkem##bits##_encap_external_entropy)                     \
+    TRAIT_METHOD(Encap, MLKEM##bits##_encap)                                   \
+    TRAIT_METHOD(Decap, MLKEM##bits##_decap)                                   \
+  };
+
+MAKE_MLKEM_TRAITS(768)
+MAKE_MLKEM_TRAITS(1024)
+
 template <typename T>
 std::vector<uint8_t> Marshal(int (*marshal_func)(CBB *, const T *),
                              const T *t) {
@@ -51,90 +88,39 @@
   return ret;
 }
 
-// These functions wrap the methods that are only in the BCM interface. They
-// take care of casting the key types from the public keys to the BCM types.
-// That saves casting noise in the template functions.
+template <typename T>
+std::vector<uint8_t> Marshal(bcm_status (*marshal_func)(CBB *, const T *),
+                             const T *t) {
+  ScopedCBB cbb;
+  uint8_t *encoded;
+  size_t encoded_len;
+  if (!CBB_init(cbb.get(), 1) ||                   //
+      !bcm_success(marshal_func(cbb.get(), t)) ||  //
+      !CBB_finish(cbb.get(), &encoded, &encoded_len)) {
+    abort();
+  }
 
-int wrapper_768_marshal_private_key(CBB *out,
-                                    const MLKEM768_private_key *private_key) {
-  return bcm_success(BCM_mlkem768_marshal_private_key(out, private_key));
+  std::vector<uint8_t> ret(encoded, encoded + encoded_len);
+  OPENSSL_free(encoded);
+  return ret;
 }
 
-int wrapper_1024_marshal_private_key(CBB *out,
-                                     const MLKEM1024_private_key *private_key) {
-  return bcm_success(BCM_mlkem1024_marshal_private_key(out, private_key));
-}
-
-void wrapper_768_generate_key_external_seed(
-    uint8_t out_encoded_public_key[MLKEM768_PUBLIC_KEY_BYTES],
-    struct MLKEM768_private_key *out_private_key,
-    const uint8_t seed[MLKEM_SEED_BYTES]) {
-  BCM_mlkem768_generate_key_external_seed(out_encoded_public_key,
-                                          out_private_key, seed);
-}
-
-void wrapper_1024_generate_key_external_seed(
-    uint8_t out_encoded_public_key[MLKEM1024_PUBLIC_KEY_BYTES],
-    struct MLKEM1024_private_key *out_private_key,
-    const uint8_t seed[MLKEM_SEED_BYTES]) {
-  BCM_mlkem1024_generate_key_external_seed(out_encoded_public_key,
-                                           out_private_key, seed);
-}
-
-void wrapper_768_encap_external_entropy(
-    uint8_t out_ciphertext[MLKEM768_CIPHERTEXT_BYTES],
-    uint8_t out_shared_secret[MLKEM_SHARED_SECRET_BYTES],
-    const MLKEM768_public_key *public_key,
-    const uint8_t entropy[BCM_MLKEM_ENCAP_ENTROPY]) {
-  BCM_mlkem768_encap_external_entropy(out_ciphertext, out_shared_secret,
-                                      public_key, entropy);
-}
-
-void wrapper_1024_encap_external_entropy(
-    uint8_t out_ciphertext[MLKEM1024_CIPHERTEXT_BYTES],
-    uint8_t out_shared_secret[MLKEM_SHARED_SECRET_BYTES],
-    const MLKEM1024_public_key *public_key,
-    const uint8_t entropy[BCM_MLKEM_ENCAP_ENTROPY]) {
-  BCM_mlkem1024_encap_external_entropy(out_ciphertext, out_shared_secret,
-                                       public_key, entropy);
-}
-
-int wrapper_768_parse_private_key(struct MLKEM768_private_key *out_private_key,
-                                  CBS *in) {
-  return bcm_success(BCM_mlkem768_parse_private_key(out_private_key, in));
-}
-
-int wrapper_1024_parse_private_key(
-    struct MLKEM1024_private_key *out_private_key, CBS *in) {
-  return bcm_success(BCM_mlkem1024_parse_private_key(out_private_key, in));
-}
-
-template <typename PUBLIC_KEY, size_t PUBLIC_KEY_BYTES, typename PRIVATE_KEY,
-          size_t PRIVATE_KEY_BYTES,
-          void (*GENERATE)(uint8_t *, uint8_t *, PRIVATE_KEY *),
-          int (*FROM_SEED)(PRIVATE_KEY *, const uint8_t *, size_t),
-          void (*PUBLIC_FROM_PRIVATE)(PUBLIC_KEY *, const PRIVATE_KEY *),
-          int (*PARSE_PUBLIC)(PUBLIC_KEY *, CBS *),
-          int (*MARSHAL_PUBLIC)(CBB *, const PUBLIC_KEY *),
-          int (*PARSE_PRIVATE)(PRIVATE_KEY *, CBS *),
-          int (*MARSHAL_PRIVATE)(CBB *, const PRIVATE_KEY *),
-          size_t CIPHERTEXT_BYTES,
-          void (*ENCAP)(uint8_t *, uint8_t *, const PUBLIC_KEY *),
-          int (*DECAP)(uint8_t *, const uint8_t *, size_t, const PRIVATE_KEY *)>
+template <typename Traits>
 void BasicTest() {
   // This function makes several ML-KEM keys, which runs up against stack
   // limits. Heap-allocate them instead.
 
-  uint8_t encoded_public_key[PUBLIC_KEY_BYTES];
+  uint8_t encoded_public_key[Traits::kPublicKeyBytes];
   uint8_t seed[MLKEM_SEED_BYTES];
-  auto priv = std::make_unique<PRIVATE_KEY>();
-  GENERATE(encoded_public_key, seed, priv.get());
+  auto priv = std::make_unique<typename Traits::PrivateKey>();
+  Traits::GenerateKey(encoded_public_key, seed, priv.get());
 
   {
-    auto priv2 = std::make_unique<PRIVATE_KEY>();
-    ASSERT_TRUE(FROM_SEED(priv2.get(), seed, sizeof(seed)));
-    EXPECT_EQ(Bytes(Declassified(Marshal(MARSHAL_PRIVATE, priv.get()))),
-              Bytes(Declassified(Marshal(MARSHAL_PRIVATE, priv2.get()))));
+    auto priv2 = std::make_unique<typename Traits::PrivateKey>();
+    ASSERT_TRUE(Traits::PrivateKeyFromSeed(priv2.get(), seed, sizeof(seed)));
+    EXPECT_EQ(
+        Bytes(Declassified(Marshal(Traits::MarshalPrivateKey, priv.get()))),
+        Bytes(Declassified(Marshal(Traits::MarshalPrivateKey, priv2.get()))));
   }
 
   uint8_t first_two_bytes[2];
@@ -143,27 +129,27 @@
   CBS encoded_public_key_cbs;
   CBS_init(&encoded_public_key_cbs, encoded_public_key,
            sizeof(encoded_public_key));
-  auto pub = std::make_unique<PUBLIC_KEY>();
+  auto pub = std::make_unique<typename Traits::PublicKey>();
   // Parsing should fail because the first coefficient is >= kPrime;
-  ASSERT_FALSE(PARSE_PUBLIC(pub.get(), &encoded_public_key_cbs));
+  ASSERT_FALSE(Traits::ParsePublicKey(pub.get(), &encoded_public_key_cbs));
 
   OPENSSL_memcpy(encoded_public_key, first_two_bytes, sizeof(first_two_bytes));
   CBS_init(&encoded_public_key_cbs, encoded_public_key,
            sizeof(encoded_public_key));
-  ASSERT_TRUE(PARSE_PUBLIC(pub.get(), &encoded_public_key_cbs));
+  ASSERT_TRUE(Traits::ParsePublicKey(pub.get(), &encoded_public_key_cbs));
   EXPECT_EQ(CBS_len(&encoded_public_key_cbs), 0u);
 
   EXPECT_EQ(Bytes(encoded_public_key),
-            Bytes(Marshal(MARSHAL_PUBLIC, pub.get())));
+            Bytes(Marshal(Traits::MarshalPublicKey, pub.get())));
 
-  auto pub2 = std::make_unique<PUBLIC_KEY>();
-  PUBLIC_FROM_PRIVATE(pub2.get(), priv.get());
+  auto pub2 = std::make_unique<typename Traits::PublicKey>();
+  Traits::PublicFromPrivate(pub2.get(), priv.get());
   EXPECT_EQ(Bytes(encoded_public_key),
-            Bytes(Marshal(MARSHAL_PUBLIC, pub2.get())));
+            Bytes(Marshal(Traits::MarshalPublicKey, pub2.get())));
 
   std::vector<uint8_t> encoded_private_key(
-      Marshal(MARSHAL_PRIVATE, priv.get()));
-  EXPECT_EQ(encoded_private_key.size(), size_t{PRIVATE_KEY_BYTES});
+      Marshal(Traits::MarshalPrivateKey, priv.get()));
+  EXPECT_EQ(encoded_private_key.size(), size_t{Traits::kPrivateKeyBytes});
 
   // Parsing should fail if a coefficient is out of range.
   {
@@ -172,8 +158,8 @@
     invalid[1] = 0xff;
     CBS cbs;
     CBS_init(&cbs, invalid.data(), invalid.size());
-    auto priv2 = std::make_unique<PRIVATE_KEY>();
-    ASSERT_FALSE(PARSE_PRIVATE(priv2.get(), &cbs));
+    auto priv2 = std::make_unique<typename Traits::PrivateKey>();
+    ASSERT_FALSE(bcm_success(Traits::ParsePrivateKey(priv2.get(), &cbs)));
   }
 
   // Parsing should fail if the public key hash is wrong.
@@ -184,54 +170,36 @@
     invalid[invalid.size() - 33] ^= 1;
     CBS cbs;
     CBS_init(&cbs, invalid.data(), invalid.size());
-    auto priv2 = std::make_unique<PRIVATE_KEY>();
-    ASSERT_FALSE(PARSE_PRIVATE(priv2.get(), &cbs));
+    auto priv2 = std::make_unique<typename Traits::PrivateKey>();
+    ASSERT_FALSE(bcm_success(Traits::ParsePrivateKey(priv2.get(), &cbs)));
   }
 
   CBS cbs;
   CBS_init(&cbs, encoded_private_key.data(), encoded_private_key.size());
-  auto priv2 = std::make_unique<PRIVATE_KEY>();
-  ASSERT_TRUE(PARSE_PRIVATE(priv2.get(), &cbs));
-  EXPECT_EQ(Bytes(Declassified(encoded_private_key)),
-            Bytes(Declassified(Marshal(MARSHAL_PRIVATE, priv2.get()))));
+  auto priv2 = std::make_unique<typename Traits::PrivateKey>();
+  ASSERT_TRUE(bcm_success(Traits::ParsePrivateKey(priv2.get(), &cbs)));
+  EXPECT_EQ(
+      Bytes(Declassified(encoded_private_key)),
+      Bytes(Declassified(Marshal(Traits::MarshalPrivateKey, priv2.get()))));
 
-  uint8_t ciphertext[CIPHERTEXT_BYTES];
+  uint8_t ciphertext[Traits::kCiphertextBytes];
   uint8_t shared_secret1[MLKEM_SHARED_SECRET_BYTES];
   uint8_t shared_secret2[MLKEM_SHARED_SECRET_BYTES];
-  ENCAP(ciphertext, shared_secret1, pub.get());
-  ASSERT_TRUE(
-      DECAP(shared_secret2, ciphertext, sizeof(ciphertext), priv.get()));
+  Traits::Encap(ciphertext, shared_secret1, pub.get());
+  ASSERT_TRUE(Traits::Decap(shared_secret2, ciphertext, sizeof(ciphertext),
+                            priv.get()));
   EXPECT_EQ(Bytes(Declassified(shared_secret1)),
             Bytes(Declassified(shared_secret2)));
-  ASSERT_TRUE(
-      DECAP(shared_secret2, ciphertext, sizeof(ciphertext), priv2.get()));
+  ASSERT_TRUE(Traits::Decap(shared_secret2, ciphertext, sizeof(ciphertext),
+                            priv2.get()));
   EXPECT_EQ(Bytes(Declassified(shared_secret1)),
             Bytes(Declassified(shared_secret2)));
 }
 
-TEST(MLKEMTest, Basic768) {
-  BasicTest<MLKEM768_public_key, MLKEM768_PUBLIC_KEY_BYTES,
-            MLKEM768_private_key, BCM_MLKEM768_PRIVATE_KEY_BYTES,
-            MLKEM768_generate_key, MLKEM768_private_key_from_seed,
-            MLKEM768_public_from_private, MLKEM768_parse_public_key,
-            MLKEM768_marshal_public_key, wrapper_768_parse_private_key,
-            wrapper_768_marshal_private_key, MLKEM768_CIPHERTEXT_BYTES,
-            MLKEM768_encap, MLKEM768_decap>();
-}
+TEST(MLKEMTest, Basic768) { BasicTest<MLKEM768Traits>(); }
+TEST(MLKEMTest, Basic1024) { BasicTest<MLKEM1024Traits>(); }
 
-TEST(MLKEMTest, Basic1024) {
-  BasicTest<MLKEM1024_public_key, MLKEM1024_PUBLIC_KEY_BYTES,
-            MLKEM1024_private_key, BCM_MLKEM1024_PRIVATE_KEY_BYTES,
-            MLKEM1024_generate_key, MLKEM1024_private_key_from_seed,
-            MLKEM1024_public_from_private, MLKEM1024_parse_public_key,
-            MLKEM1024_marshal_public_key, wrapper_1024_parse_private_key,
-            wrapper_1024_marshal_private_key, MLKEM1024_CIPHERTEXT_BYTES,
-            MLKEM1024_encap, MLKEM1024_decap>();
-}
-
-template <typename PUBLIC_KEY, size_t PUBLIC_KEY_BYTES, typename PRIVATE_KEY,
-          int (*MARSHAL_PRIVATE)(CBB *, const PRIVATE_KEY *),
-          void (*GENERATE)(uint8_t *, PRIVATE_KEY *, const uint8_t *)>
+template <typename Traits>
 void MLKEMKeyGenFileTest(FileTest *t) {
   std::vector<uint8_t> expected_pub_key_bytes, seed, expected_priv_key_bytes;
   ASSERT_TRUE(t->GetBytes(&seed, "seed"));
@@ -241,11 +209,12 @@
 
   ASSERT_EQ(seed.size(), size_t{MLKEM_SEED_BYTES});
 
-  std::vector<uint8_t> pub_key_bytes(PUBLIC_KEY_BYTES);
-  auto priv = std::make_unique<PRIVATE_KEY>();
-  GENERATE(pub_key_bytes.data(), priv.get(), seed.data());
+  std::vector<uint8_t> pub_key_bytes(Traits::kPublicKeyBytes);
+  auto priv = std::make_unique<typename Traits::PrivateKey>();
+  Traits::GenerateKeyExternalSeed(pub_key_bytes.data(), priv.get(),
+                                  seed.data());
   const std::vector<uint8_t> priv_key_bytes(
-      Marshal(MARSHAL_PRIVATE, priv.get()));
+      Marshal(Traits::MarshalPrivateKey, priv.get()));
 
   EXPECT_EQ(Bytes(pub_key_bytes), Bytes(expected_pub_key_bytes));
   EXPECT_EQ(Bytes(Declassified(priv_key_bytes)),
@@ -253,25 +222,16 @@
 }
 
 TEST(MLKEMTest, KeyGen768TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem768_keygen_tests.txt",
-      MLKEMKeyGenFileTest<MLKEM768_public_key, MLKEM768_PUBLIC_KEY_BYTES,
-                          MLKEM768_private_key, wrapper_768_marshal_private_key,
-                          wrapper_768_generate_key_external_seed>);
+  FileTestGTest("crypto/mlkem/mlkem768_keygen_tests.txt",
+                MLKEMKeyGenFileTest<MLKEM768Traits>);
 }
 
 TEST(MLKEMTest, KeyGen1024TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem1024_keygen_tests.txt",
-      MLKEMKeyGenFileTest<MLKEM1024_public_key, MLKEM1024_PUBLIC_KEY_BYTES,
-                          MLKEM1024_private_key,
-                          wrapper_1024_marshal_private_key,
-                          wrapper_1024_generate_key_external_seed>);
+  FileTestGTest("crypto/mlkem/mlkem1024_keygen_tests.txt",
+                MLKEMKeyGenFileTest<MLKEM1024Traits>);
 }
 
-template <typename PUBLIC_KEY, size_t PUBLIC_KEY_BYTES, typename PRIVATE_KEY,
-          int (*MARSHAL_PRIVATE)(CBB *, const PRIVATE_KEY *),
-          void (*GENERATE)(uint8_t *, PRIVATE_KEY *, const uint8_t *)>
+template <typename Traits>
 void MLKEMNistKeyGenFileTest(FileTest *t) {
   std::vector<uint8_t> expected_pub_key_bytes, z, d, expected_priv_key_bytes;
   ASSERT_TRUE(t->GetBytes(&z, "z"));
@@ -285,38 +245,27 @@
   uint8_t seed[MLKEM_SEED_BYTES];
   OPENSSL_memcpy(&seed[0], d.data(), d.size());
   OPENSSL_memcpy(&seed[MLKEM_SEED_BYTES / 2], z.data(), z.size());
-  std::vector<uint8_t> pub_key_bytes(PUBLIC_KEY_BYTES);
-  auto priv = std::make_unique<PRIVATE_KEY>();
-  GENERATE(pub_key_bytes.data(), priv.get(), seed);
+  std::vector<uint8_t> pub_key_bytes(Traits::kPublicKeyBytes);
+  auto priv = std::make_unique<typename Traits::PrivateKey>();
+  Traits::GenerateKeyExternalSeed(pub_key_bytes.data(), priv.get(), seed);
   const std::vector<uint8_t> priv_key_bytes(
-      Marshal(MARSHAL_PRIVATE, priv.get()));
+      Marshal(Traits::MarshalPrivateKey, priv.get()));
 
   EXPECT_EQ(Bytes(pub_key_bytes), Bytes(expected_pub_key_bytes));
   EXPECT_EQ(Bytes(priv_key_bytes), Bytes(expected_priv_key_bytes));
 }
 
 TEST(MLKEMTest, NISTKeyGen768TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem768_nist_keygen_tests.txt",
-      MLKEMNistKeyGenFileTest<MLKEM768_public_key, MLKEM768_PUBLIC_KEY_BYTES,
-                              MLKEM768_private_key,
-                              wrapper_768_marshal_private_key,
-                              wrapper_768_generate_key_external_seed>);
+  FileTestGTest("crypto/mlkem/mlkem768_nist_keygen_tests.txt",
+                MLKEMNistKeyGenFileTest<MLKEM768Traits>);
 }
 
 TEST(MLKEMTest, NISTKeyGen1024TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem1024_nist_keygen_tests.txt",
-      MLKEMNistKeyGenFileTest<MLKEM1024_public_key, MLKEM1024_PUBLIC_KEY_BYTES,
-                              MLKEM1024_private_key,
-                              wrapper_1024_marshal_private_key,
-                              wrapper_1024_generate_key_external_seed>);
+  FileTestGTest("crypto/mlkem/mlkem1024_nist_keygen_tests.txt",
+                MLKEMNistKeyGenFileTest<MLKEM1024Traits>);
 }
 
-template <typename PUBLIC_KEY, size_t PUBLIC_KEY_BYTES,
-          int (*PARSE_PUBLIC)(PUBLIC_KEY *, CBS *), size_t CIPHERTEXT_BYTES,
-          void (*ENCAP)(uint8_t *, uint8_t *, const PUBLIC_KEY *,
-                        const uint8_t *)>
+template <typename Traits>
 void MLKEMEncapFileTest(FileTest *t) {
   std::vector<uint8_t> pub_key_bytes, entropy, expected_ciphertext,
       expected_shared_secret;
@@ -328,42 +277,35 @@
   std::string result;
   ASSERT_TRUE(t->GetAttribute(&result, "result"));
 
-  PUBLIC_KEY pub_key;
+  typename Traits::PublicKey pub_key;
   CBS pub_key_cbs;
   CBS_init(&pub_key_cbs, pub_key_bytes.data(), pub_key_bytes.size());
-  const int parse_ok = PARSE_PUBLIC(&pub_key, &pub_key_cbs);
+  const int parse_ok = Traits::ParsePublicKey(&pub_key, &pub_key_cbs);
   ASSERT_EQ(parse_ok, result == "pass");
   if (!parse_ok) {
     return;
   }
 
-  uint8_t ciphertext[CIPHERTEXT_BYTES];
+  uint8_t ciphertext[Traits::kCiphertextBytes];
   uint8_t shared_secret[MLKEM_SHARED_SECRET_BYTES];
-  ENCAP(ciphertext, shared_secret, &pub_key, entropy.data());
+  Traits::EncapExternalEntropy(ciphertext, shared_secret, &pub_key,
+                               entropy.data());
 
   ASSERT_EQ(Bytes(expected_ciphertext), Bytes(ciphertext));
   ASSERT_EQ(Bytes(expected_shared_secret), Bytes(Declassified(shared_secret)));
 }
 
 TEST(MLKEMTest, Encap768TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem768_encap_tests.txt",
-      MLKEMEncapFileTest<MLKEM768_public_key, MLKEM768_PUBLIC_KEY_BYTES,
-                         MLKEM768_parse_public_key, MLKEM768_CIPHERTEXT_BYTES,
-                         wrapper_768_encap_external_entropy>);
+  FileTestGTest("crypto/mlkem/mlkem768_encap_tests.txt",
+                MLKEMEncapFileTest<MLKEM768Traits>);
 }
 
 TEST(MLKEMTest, Encap1024TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem1024_encap_tests.txt",
-      MLKEMEncapFileTest<MLKEM1024_public_key, MLKEM1024_PUBLIC_KEY_BYTES,
-                         MLKEM1024_parse_public_key, MLKEM1024_CIPHERTEXT_BYTES,
-                         wrapper_1024_encap_external_entropy>);
+  FileTestGTest("crypto/mlkem/mlkem1024_encap_tests.txt",
+                MLKEMEncapFileTest<MLKEM1024Traits>);
 }
 
-template <typename PRIVATE_KEY, size_t PRIVATE_KEY_BYTES,
-          int (*PARSE_PRIVATE)(PRIVATE_KEY *, CBS *), size_t CIPHERTEXT_BYTES,
-          int (*DECAP)(uint8_t *, const uint8_t *, size_t, const PRIVATE_KEY *)>
+template <typename Traits>
 void MLKEMDecapFileTest(FileTest *t) {
   std::vector<uint8_t> priv_key_bytes, ciphertext, expected_shared_secret;
   ASSERT_TRUE(t->GetBytes(&priv_key_bytes, "private_key"));
@@ -372,18 +314,19 @@
   std::string result;
   ASSERT_TRUE(t->GetAttribute(&result, "result"));
 
-  PRIVATE_KEY priv_key;
+  typename Traits::PrivateKey priv_key;
   CBS priv_key_cbs;
   CBS_init(&priv_key_cbs, priv_key_bytes.data(), priv_key_bytes.size());
-  const int parse_ok = PARSE_PRIVATE(&priv_key, &priv_key_cbs);
+  const int parse_ok =
+      bcm_success(Traits::ParsePrivateKey(&priv_key, &priv_key_cbs));
   if (!parse_ok) {
     ASSERT_NE(result, "pass");
     return;
   }
 
   uint8_t shared_secret[MLKEM_SHARED_SECRET_BYTES];
-  const int decap_ok =
-      DECAP(shared_secret, ciphertext.data(), ciphertext.size(), &priv_key);
+  const int decap_ok = Traits::Decap(shared_secret, ciphertext.data(),
+                                     ciphertext.size(), &priv_key);
   if (!decap_ok) {
     ASSERT_NE(result, "pass");
     return;
@@ -393,109 +336,92 @@
 }
 
 TEST(MLKEMTest, Decap768TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem768_decap_tests.txt",
-      MLKEMDecapFileTest<MLKEM768_private_key, BCM_MLKEM768_PRIVATE_KEY_BYTES,
-                         wrapper_768_parse_private_key,
-                         MLKEM768_CIPHERTEXT_BYTES, MLKEM768_decap>);
+  FileTestGTest("crypto/mlkem/mlkem768_decap_tests.txt",
+                MLKEMDecapFileTest<MLKEM768Traits>);
 }
 
 TEST(MLKEMTest, Decap1024TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem1024_decap_tests.txt",
-      MLKEMDecapFileTest<MLKEM1024_private_key, BCM_MLKEM1024_PRIVATE_KEY_BYTES,
-                         wrapper_1024_parse_private_key,
-                         MLKEM1024_CIPHERTEXT_BYTES, MLKEM1024_decap>);
+  FileTestGTest("crypto/mlkem/mlkem1024_decap_tests.txt",
+                MLKEMDecapFileTest<MLKEM1024Traits>);
 }
 
-template <typename PRIVATE_KEY, int (*PARSE_PRIVATE)(PRIVATE_KEY *, CBS *),
-          int (*DECAP)(uint8_t *, const uint8_t *, size_t, const PRIVATE_KEY *)>
+template <typename Traits>
 void MLKEMNistDecapFileTest(FileTest *t) {
   std::vector<uint8_t> ciphertext, expected_shared_secret, private_key_bytes;
   ASSERT_TRUE(t->GetBytes(&ciphertext, "c"));
   ASSERT_TRUE(t->GetBytes(&expected_shared_secret, "k"));
   ASSERT_TRUE(t->GetInstructionBytes(&private_key_bytes, "dk"));
 
-  PRIVATE_KEY priv;
+  typename Traits::PrivateKey priv;
   CBS private_key_cbs;
   CBS_init(&private_key_cbs, private_key_bytes.data(),
            private_key_bytes.size());
-  ASSERT_TRUE(PARSE_PRIVATE(&priv, &private_key_cbs));
+  ASSERT_TRUE(bcm_success(Traits::ParsePrivateKey(&priv, &private_key_cbs)));
 
   uint8_t shared_secret[MLKEM_SHARED_SECRET_BYTES];
-  ASSERT_TRUE(
-      DECAP(shared_secret, ciphertext.data(), ciphertext.size(), &priv));
+  ASSERT_TRUE(Traits::Decap(shared_secret, ciphertext.data(), ciphertext.size(),
+                            &priv));
 
   ASSERT_EQ(Bytes(shared_secret), Bytes(expected_shared_secret));
 }
 
 TEST(MLKEMTest, NistDecap768TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem768_nist_decap_tests.txt",
-      MLKEMNistDecapFileTest<MLKEM768_private_key,
-                             wrapper_768_parse_private_key, MLKEM768_decap>);
+  FileTestGTest("crypto/mlkem/mlkem768_nist_decap_tests.txt",
+                MLKEMNistDecapFileTest<MLKEM768Traits>);
 }
 
 TEST(MLKEMTest, NistDecap1024TestVectors) {
-  FileTestGTest(
-      "crypto/mlkem/mlkem1024_nist_decap_tests.txt",
-      MLKEMNistDecapFileTest<MLKEM1024_private_key,
-                             wrapper_1024_parse_private_key, MLKEM1024_decap>);
+  FileTestGTest("crypto/mlkem/mlkem1024_nist_decap_tests.txt",
+                MLKEMNistDecapFileTest<MLKEM1024Traits>);
 }
 
 // Unoptimized builds are much slower, and iterative tests run ML-KEM many
 // times. Disable them in unoptimized builds for now.
 // https://crbug.com/479850443
 #if (defined(__GNUC__) || defined(__clang__)) && !defined(__OPTIMIZE__)
-#define DISABLE_IF_NOT_OPTIMIZED(t) DISABLED_ ## t
+#define DISABLE_IF_NOT_OPTIMIZED(t) DISABLED_##t
 #else
 #define DISABLE_IF_NOT_OPTIMIZED(t) t
 #endif
 
-template <
-    typename PUBLIC_KEY, size_t PUBLIC_KEY_BYTES, typename PRIVATE_KEY,
-    size_t PRIVATE_KEY_BYTES,
-    void (*GENERATE)(uint8_t *, PRIVATE_KEY *, const uint8_t *),
-    void (*TO_PUBLIC)(PUBLIC_KEY *, const PRIVATE_KEY *),
-    int (*MARSHAL_PRIVATE)(CBB *, const PRIVATE_KEY *), size_t CIPHERTEXT_BYTES,
-    void (*ENCAP)(uint8_t *, uint8_t *, const PUBLIC_KEY *, const uint8_t *),
-    int (*DECAP)(uint8_t *, const uint8_t *, size_t, const PRIVATE_KEY *)>
+template <typename Traits>
 void IteratedTest(uint8_t out[32]) {
   BORINGSSL_keccak_st generate_st;
   BORINGSSL_keccak_init(&generate_st, boringssl_shake128);
   BORINGSSL_keccak_st results_st;
   BORINGSSL_keccak_init(&results_st, boringssl_shake128);
 
-  auto priv = std::make_unique<PRIVATE_KEY>();
-  auto pub = std::make_unique<PUBLIC_KEY>();
+  auto priv = std::make_unique<typename Traits::PrivateKey>();
+  auto pub = std::make_unique<typename Traits::PublicKey>();
   for (int i = 0; i < 10000; i++) {
     uint8_t seed[MLKEM_SEED_BYTES];
     BORINGSSL_keccak_squeeze(&generate_st, seed, sizeof(seed));
-    uint8_t encoded_pub[PUBLIC_KEY_BYTES];
-    GENERATE(encoded_pub, priv.get(), seed);
-    TO_PUBLIC(pub.get(), priv.get());
+    uint8_t encoded_pub[Traits::kPublicKeyBytes];
+    Traits::GenerateKeyExternalSeed(encoded_pub, priv.get(), seed);
+    Traits::PublicFromPrivate(pub.get(), priv.get());
 
     BORINGSSL_keccak_absorb(&results_st, encoded_pub, sizeof(encoded_pub));
     const std::vector<uint8_t> encoded_priv(
-        Marshal(MARSHAL_PRIVATE, priv.get()));
+        Marshal(Traits::MarshalPrivateKey, priv.get()));
     BORINGSSL_keccak_absorb(&results_st, encoded_priv.data(),
                             encoded_priv.size());
 
     uint8_t encap_entropy[BCM_MLKEM_ENCAP_ENTROPY];
     BORINGSSL_keccak_squeeze(&generate_st, encap_entropy,
                              sizeof(encap_entropy));
-    uint8_t ciphertext[CIPHERTEXT_BYTES];
+    uint8_t ciphertext[Traits::kCiphertextBytes];
     uint8_t shared_secret[MLKEM_SHARED_SECRET_BYTES];
-    ENCAP(ciphertext, shared_secret, pub.get(), encap_entropy);
+    Traits::EncapExternalEntropy(ciphertext, shared_secret, pub.get(),
+                                 encap_entropy);
 
     BORINGSSL_keccak_absorb(&results_st, ciphertext, sizeof(ciphertext));
     BORINGSSL_keccak_absorb(&results_st, shared_secret, sizeof(shared_secret));
 
-    uint8_t invalid_ciphertext[CIPHERTEXT_BYTES];
+    uint8_t invalid_ciphertext[Traits::kCiphertextBytes];
     BORINGSSL_keccak_squeeze(&generate_st, invalid_ciphertext,
                              sizeof(invalid_ciphertext));
-    ASSERT_TRUE(DECAP(shared_secret, invalid_ciphertext,
-                      sizeof(invalid_ciphertext), priv.get()));
+    ASSERT_TRUE(Traits::Decap(shared_secret, invalid_ciphertext,
+                              sizeof(invalid_ciphertext), priv.get()));
 
     BORINGSSL_keccak_absorb(&results_st, shared_secret, sizeof(shared_secret));
   }
@@ -509,12 +435,7 @@
   // but the final value has been updated to reflect the change from Kyber to
   // ML-KEM.
   uint8_t result[32];
-  IteratedTest<MLKEM768_public_key, MLKEM768_PUBLIC_KEY_BYTES,
-               MLKEM768_private_key, BCM_MLKEM768_PRIVATE_KEY_BYTES,
-               wrapper_768_generate_key_external_seed,
-               MLKEM768_public_from_private, wrapper_768_marshal_private_key,
-               MLKEM768_CIPHERTEXT_BYTES, wrapper_768_encap_external_entropy,
-               MLKEM768_decap>(result);
+  IteratedTest<MLKEM768Traits>(result);
 
   const uint8_t kExpected[32] = {
       0xf9, 0x59, 0xd1, 0x8d, 0x3d, 0x11, 0x80, 0x12, 0x14, 0x33, 0xbf,
@@ -530,12 +451,7 @@
   // but the final value has been updated to reflect the change from Kyber to
   // ML-KEM.
   uint8_t result[32];
-  IteratedTest<MLKEM1024_public_key, MLKEM1024_PUBLIC_KEY_BYTES,
-               MLKEM1024_private_key, BCM_MLKEM1024_PRIVATE_KEY_BYTES,
-               wrapper_1024_generate_key_external_seed,
-               MLKEM1024_public_from_private, wrapper_1024_marshal_private_key,
-               MLKEM1024_CIPHERTEXT_BYTES, wrapper_1024_encap_external_entropy,
-               MLKEM1024_decap>(result);
+  IteratedTest<MLKEM1024Traits>(result);
 
   const uint8_t kExpected[32] = {
       0xe3, 0xbf, 0x82, 0xb0, 0x13, 0x30, 0x7b, 0x2e, 0x9d, 0x47, 0xdd,