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,