From 79fd85d1e8b055b93dbb796d181c073b5973d628 Mon Sep 17 00:00:00 2001 From: Gilles Peskine Date: Tue, 28 Jul 2026 19:50:12 +0200 Subject: [PATCH] Clean up how private key formats are specified to DriverGenerator Instead of having `gen_all()` change the object state, have the constructor set the object field as desired. This doesn't change the runtime behavior since we're only running `gen_all()` once after constructing the object, and nothing else. But do clean thing up, to avoid any surprises later. Signed-off-by: Gilles Peskine --- util/mbedtls_maintainer/mldsa_test_generator.py | 16 +++++++++------- 1 file changed, 9 insertions(+), 7 deletions(-) diff --git a/util/mbedtls_maintainer/mldsa_test_generator.py b/util/mbedtls_maintainer/mldsa_test_generator.py index 812aa649b..09308fc58 100644 --- a/util/mbedtls_maintainer/mldsa_test_generator.py +++ b/util/mbedtls_maintainer/mldsa_test_generator.py @@ -94,8 +94,14 @@ MESSAGES = [ class Generator: """Abstract base class to generate tests for one API.""" - def __init__(self) -> None: - self.private_key_formats: Sequence[PrivateKeyFormat] = [PrivateKeyFormat.SEED] + PRIVATE_KEY_FORMATS: Sequence[PrivateKeyFormat] = [PrivateKeyFormat.SEED] + + def __init__(self, + private_key_formats: Optional[Sequence[PrivateKeyFormat]] = None, + ) -> None: + self.private_key_formats = \ + private_key_formats if private_key_formats is not None else \ + self.PRIVATE_KEY_FORMATS @classmethod def function(cls, func: str, kl: int) -> str: @@ -243,8 +249,7 @@ class Generator: class PQCPGenerator(Generator): """Test mldsa-native entry points.""" - def __init__(self) -> None: - self.private_key_formats = [PrivateKeyFormat.EXPANDED] + PRIVATE_KEY_FORMATS = [PrivateKeyFormat.EXPANDED] @classmethod def function(cls, func: str, kl: int) -> str: @@ -401,11 +406,8 @@ class DriverGenerator(Generator): def gen_all(self, multipart: bool = False, - private_key_formats: Optional[Sequence[PrivateKeyFormat]] = None, ) -> Iterator[test_case.TestCase]: """Generate all the tests for this API.""" - if private_key_formats is not None: - self.private_key_formats = private_key_formats yield from super().gen_all() if multipart: for kl in sorted(KEYS.keys()):