ACVP: add extra ML-KEM tests.

These appear to be effectively required in practice.

Change-Id: I7d2fe54454f0cd65c55f9ca1a1b305af50508850
Reviewed-on: https://boringssl-review.googlesource.com/c/boringssl/+/97827
Auto-Submit: Adam Langley <agl@google.com>
Reviewed-by: David Benjamin <davidben@google.com>
SLSA-Policy-Verified: SLSA Policy Verification Service <devtools-gerritcodereview-exitgate@google.com>
Commit-Queue: Adam Langley <agl@google.com>
diff --git a/util/fipstools/acvp/acvptool/subprocess/mlkem.go b/util/fipstools/acvp/acvptool/subprocess/mlkem.go
index a112ed8..883fd48 100644
--- a/util/fipstools/acvp/acvptool/subprocess/mlkem.go
+++ b/util/fipstools/acvp/acvptool/subprocess/mlkem.go
@@ -76,9 +76,10 @@
 }
 
 type mlkemEncapDecapTestResponse struct {
-	ID uint64 `json:"tcId"`
-	C  string `json:"c,omitempty"`
-	K  string `json:"k,omitempty"`
+	ID         uint64 `json:"tcId"`
+	C          string `json:"c,omitempty"`
+	K          string `json:"k,omitempty"`
+	TestPassed *bool  `json:"testPassed,omitempty"`
 }
 
 func decodeNonEmptyHex(in string) ([]byte, error) {
@@ -238,6 +239,50 @@
 				})
 			}
 
+		case "encapsulationKeyCheck":
+			cmdName := group.ParameterSet + "/encapKeyCheck"
+			for _, test := range group.Tests {
+				ek, err := decodeNonEmptyHex(test.EK)
+				if err != nil {
+					return nil, fmt.Errorf("failed to decode ek in test case %d/%d: %s",
+						group.ID, test.ID, err)
+				}
+
+				result, err := t.Transact(cmdName, 1, ek)
+				if err != nil {
+					return nil, fmt.Errorf("encapsulation key check failed for test case %d/%d: %s",
+						group.ID, test.ID, err)
+				}
+
+				passed := result[0][0] != 0
+				response.Tests = append(response.Tests, mlkemEncapDecapTestResponse{
+					ID:         test.ID,
+					TestPassed: &passed,
+				})
+			}
+
+		case "decapsulationKeyCheck":
+			cmdName := group.ParameterSet + "/decapKeyCheck"
+			for _, test := range group.Tests {
+				dk, err := decodeNonEmptyHex(test.DK)
+				if err != nil {
+					return nil, fmt.Errorf("failed to decode dk in test case %d/%d: %s",
+						group.ID, test.ID, err)
+				}
+
+				result, err := t.Transact(cmdName, 1, dk)
+				if err != nil {
+					return nil, fmt.Errorf("decapsulation key check failed for test case %d/%d: %s",
+						group.ID, test.ID, err)
+				}
+
+				passed := result[0][0] != 0
+				response.Tests = append(response.Tests, mlkemEncapDecapTestResponse{
+					ID:         test.ID,
+					TestPassed: &passed,
+				})
+			}
+
 		default:
 			return nil, fmt.Errorf("unsupported function: %s", group.Function)
 		}
diff --git a/util/fipstools/acvp/acvptool/test/expected/ML-KEM.bz2 b/util/fipstools/acvp/acvptool/test/expected/ML-KEM.bz2
index 8c9ee96..1d1309c 100644
--- a/util/fipstools/acvp/acvptool/test/expected/ML-KEM.bz2
+++ b/util/fipstools/acvp/acvptool/test/expected/ML-KEM.bz2
Binary files differ
diff --git a/util/fipstools/acvp/acvptool/test/vectors/ML-KEM.bz2 b/util/fipstools/acvp/acvptool/test/vectors/ML-KEM.bz2
index d42032e..a454f8c 100644
--- a/util/fipstools/acvp/acvptool/test/vectors/ML-KEM.bz2
+++ b/util/fipstools/acvp/acvptool/test/vectors/ML-KEM.bz2
Binary files differ
diff --git a/util/fipstools/acvp/modulewrapper/modulewrapper.cc b/util/fipstools/acvp/modulewrapper/modulewrapper.cc
index de14f8a..ec0d279 100644
--- a/util/fipstools/acvp/modulewrapper/modulewrapper.cc
+++ b/util/fipstools/acvp/modulewrapper/modulewrapper.cc
@@ -812,7 +812,9 @@
         ],
         "functions": [
           "encapsulation",
-          "decapsulation"
+          "decapsulation",
+          "encapsulationKeyCheck",
+          "decapsulationKeyCheck"
         ]
       },
       {
@@ -1867,9 +1869,9 @@
     }
     memcpy(&salt_len, args[2].data(), sizeof(salt_len));
     if (salt_len != digest_len) {
-      LOG_ERROR(
-          "PSS salt length %u does not match digest length %u.\n",
-          static_cast<unsigned>(salt_len), static_cast<unsigned>(digest_len));
+      LOG_ERROR("PSS salt length %u does not match digest length %u.\n",
+                static_cast<unsigned>(salt_len),
+                static_cast<unsigned>(digest_len));
       return false;
     }
     if (!RSA_sign_pss_mgf1(key, &sig_len, sig.data(), sig.size(), digest_buf,
@@ -2277,6 +2279,30 @@
   return write_reply({shared_secret});
 }
 
+template <typename PublicKey, bcm_status (*ParsePublic)(PublicKey *, CBS *)>
+static bool MLKEMEncapKeyCheck(const Span<const uint8_t> args[],
+                               ReplyCallback write_reply) {
+  const Span<const uint8_t> pub_key_bytes = args[0];
+
+  auto pub = std::make_unique<PublicKey>();
+  CBS cbs = pub_key_bytes;
+  uint8_t valid = bcm_success(ParsePublic(pub.get(), &cbs));
+
+  return write_reply({Span<const uint8_t>(&valid, sizeof(valid))});
+}
+
+template <typename PrivateKey, bcm_status (*ParsePrivate)(PrivateKey *, CBS *)>
+static bool MLKEMDecapKeyCheck(const Span<const uint8_t> args[],
+                               ReplyCallback write_reply) {
+  const Span<const uint8_t> priv_key_bytes = args[0];
+
+  auto priv = std::make_unique<PrivateKey>();
+  CBS cbs = priv_key_bytes;
+  uint8_t valid = bcm_success(ParsePrivate(priv.get(), &cbs));
+
+  return write_reply({Span<const uint8_t>(&valid, sizeof(valid))});
+}
+
 template <size_t N, size_t PublicKeyBytes, size_t PrivateKeyBytes,
           bcm_infallible (*GenerateFromSeed)(uint8_t *, uint8_t *,
                                              const uint8_t *)>
@@ -2503,6 +2529,15 @@
     {"ML-KEM-1024/decap", 2,
      MLKEMDecap<MLKEM1024_private_key, BCM_mlkem1024_parse_private_key,
                 BCM_mlkem1024_decap>},
+    {"ML-KEM-768/encapKeyCheck", 1,
+     MLKEMEncapKeyCheck<MLKEM768_public_key, BCM_mlkem768_parse_public_key>},
+    {"ML-KEM-1024/encapKeyCheck", 1,
+     MLKEMEncapKeyCheck<MLKEM1024_public_key, BCM_mlkem1024_parse_public_key>},
+    {"ML-KEM-768/decapKeyCheck", 1,
+     MLKEMDecapKeyCheck<MLKEM768_private_key, BCM_mlkem768_parse_private_key>},
+    {"ML-KEM-1024/decapKeyCheck", 1,
+     MLKEMDecapKeyCheck<MLKEM1024_private_key,
+                        BCM_mlkem1024_parse_private_key>},
     {"SLH-DSA-SHA2-128s/keyGen", 1,
      SLHDSAKeyGen<BCM_SLHDSA_SHA2_128S_N, BCM_SLHDSA_SHA2_128S_PUBLIC_KEY_BYTES,
                   BCM_SLHDSA_SHA2_128S_PRIVATE_KEY_BYTES,