Merge to fips-20250107: ACVP: add extra ML-KEM tests.

These appear to be effectively required in practice.

See https://boringssl-review.googlesource.com/c/boringssl/+/97827
(cherry picked from commit 0b364bbf0c2e8686163f82484685f5a68c444172)
Change-Id: I7f717f5e8bb870e5e8dbfd39eb6804acde3e43cf
Reviewed-on: https://boringssl-review.googlesource.com/c/boringssl/+/100808
SLSA-Policy-Verified: SLSA Policy Verification Service <devtools-gerritcodereview-exitgate@google.com>
Reviewed-by: David Benjamin <davidben@google.com>
diff --git a/util/fipstools/acvp/acvptool/subprocess/mlkem.go b/util/fipstools/acvp/acvptool/subprocess/mlkem.go
index e9fd45d..883fd48 100644
--- a/util/fipstools/acvp/acvptool/subprocess/mlkem.go
+++ b/util/fipstools/acvp/acvptool/subprocess/mlkem.go
@@ -3,6 +3,7 @@
 import (
 	"encoding/hex"
 	"encoding/json"
+	"errors"
 	"fmt"
 	"strings"
 )
@@ -58,13 +59,13 @@
 	TestType     string                `json:"testType"`
 	ParameterSet string                `json:"parameterSet"`
 	Function     string                `json:"function"`
-	DK           string                `json:"dk,omitempty"`
 	Tests        []mlkemEncapDecapTest `json:"tests"`
 }
 
 type mlkemEncapDecapTest struct {
 	ID uint64 `json:"tcId"`
 	EK string `json:"ek,omitempty"`
+	DK string `json:"dk,omitempty"`
 	M  string `json:"m,omitempty"`
 	C  string `json:"c,omitempty"`
 }
@@ -75,9 +76,21 @@
 }
 
 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) {
+	ret, err := hex.DecodeString(in)
+	if err != nil {
+		return nil, err
+	}
+	if len(ret) == 0 {
+		return nil, errors.New("empty string")
+	}
+	return ret, nil
 }
 
 type mlkem struct{}
@@ -200,14 +213,15 @@
 
 		case "decapsulation":
 			cmdName := group.ParameterSet + "/decap"
-			dk, err := hex.DecodeString(group.DK)
-			if err != nil {
-				return nil, fmt.Errorf("failed to decode dk in group %d: %s",
-					group.ID, err)
-			}
 
 			for _, test := range group.Tests {
-				c, err := hex.DecodeString(test.C)
+				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)
+				}
+
+				c, err := decodeNonEmptyHex(test.C)
 				if err != nil {
 					return nil, fmt.Errorf("failed to decode c in test case %d/%d: %s",
 						group.ID, test.ID, err)
@@ -225,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..ade6a89 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 962d017..46d2860 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 06ee975..e90e510 100644
--- a/util/fipstools/acvp/modulewrapper/modulewrapper.cc
+++ b/util/fipstools/acvp/modulewrapper/modulewrapper.cc
@@ -1038,7 +1038,9 @@
         ],
         "functions": [
           "encapsulation",
-          "decapsulation"
+          "decapsulation",
+          "encapsulationKeyCheck",
+          "decapsulationKeyCheck"
         ]
       },
       {
@@ -2521,6 +2523,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))});
+}
+
 static bool SLHDSAKeyGen(const Span<const uint8_t> args[],
                          ReplyCallback write_reply) {
   const Span<const uint8_t> seed = args[0];
@@ -2728,6 +2754,18 @@
     {"ML-KEM-1024/decap", 2,
      MLKEMDecap<BCM_mlkem1024_private_key, BCM_mlkem1024_parse_private_key,
                 BCM_mlkem1024_decap>},
+    {"ML-KEM-768/encapKeyCheck", 1,
+     MLKEMEncapKeyCheck<BCM_mlkem768_public_key,
+                        BCM_mlkem768_parse_public_key>},
+    {"ML-KEM-1024/encapKeyCheck", 1,
+     MLKEMEncapKeyCheck<BCM_mlkem1024_public_key,
+                        BCM_mlkem1024_parse_public_key>},
+    {"ML-KEM-768/decapKeyCheck", 1,
+     MLKEMDecapKeyCheck<BCM_mlkem768_private_key,
+                        BCM_mlkem768_parse_private_key>},
+    {"ML-KEM-1024/decapKeyCheck", 1,
+     MLKEMDecapKeyCheck<BCM_mlkem1024_private_key,
+                        BCM_mlkem1024_parse_private_key>},
     {"SLH-DSA-SHA2-128s/keyGen", 1, SLHDSAKeyGen},
     {"SLH-DSA-SHA2-128s/sigGen", 3, SLHDSASigGen},
     {"SLH-DSA-SHA2-128s/sigVer", 3, SLHDSASigVer},