From 51cc7692c91d9eb2b552c78eb1f1ac028774ba3d Mon Sep 17 00:00:00 2001 From: Stanley Shen Date: Sat, 26 Sep 2026 20:25:19 -0700 Subject: [PATCH 1/2] Fix StatefulSignature.free leaking the native OQS_SIG_STFL struct free() detached the store callback and freed the secret key when the instance owns it, but it never called OQS_SIG_STFL_free, so the struct allocated by OQS_SIG_STFL_new in __init__ leaked for every instance, including the ones released through the context manager. KeyEncapsulation.free and Signature.free already call their OQS_KEM_free/OQS_SIG_free counterparts, and the class docstring lists free as the wrapper for OQS_SIG_STFL_free. Call OQS_SIG_STFL_free and clear the pointer, so a second free() is a no-op instead of a double free, and add a test that records the calls made to OQS_SIG_STFL_free. Co-Authored-By: Claude Opus 5 Signed-off-by: Stanley Shen --- oqs/oqs.py | 3 +++ tests/test_stfl_sig.py | 19 +++++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/oqs/oqs.py b/oqs/oqs.py index 8cad956..11a6866 100644 --- a/oqs/oqs.py +++ b/oqs/oqs.py @@ -1274,6 +1274,9 @@ def free(self) -> None: self._store_cb = None if self._secret_key and self._owns_secret: native().OQS_SIG_STFL_SECRET_KEY_free(self._secret_key) + if self._sig: + native().OQS_SIG_STFL_free(self._sig) + self._sig = None native().OQS_SIG_STFL_new.restype = ct.POINTER(StatefulSignature) diff --git a/tests/test_stfl_sig.py b/tests/test_stfl_sig.py index 6825a84..793e039 100644 --- a/tests/test_stfl_sig.py +++ b/tests/test_stfl_sig.py @@ -2,6 +2,7 @@ import platform # to learn the OS we're on import random from pathlib import Path +from unittest import mock from typing import Tuple @@ -158,6 +159,24 @@ def test_python_attributes() -> None: raise AssertionError(msg) +def test_free() -> None: + lib = oqs.native() + real_free = lib.OQS_SIG_STFL_free + freed: list[object] = [] + + def recording_free(sig_ptr: object) -> None: + freed.append(sig_ptr) + real_free(sig_ptr) + + sig = oqs.StatefulSignature(oqs.get_enabled_stateful_sig_mechanisms()[0]) + with mock.patch.object(lib, "OQS_SIG_STFL_free", recording_free): + sig.free() + # The struct must not be freed a second time. + sig.free() + + assert len(freed) == 1 # noqa: S101 + + if __name__ == "__main__": try: import nose2 From a9557504d0d4fdf64e12a4f082810a2ff4aea93c Mon Sep 17 00:00:00 2001 From: Douglas Stebila Date: Mon, 28 Sep 2026 11:33:41 -0400 Subject: [PATCH 2/2] Free the StatefulSignature secret key and guard use after free StatefulSignature now always owns its native secret key. free() releases it with OQS_SIG_STFL_SECRET_KEY_free, which securely erases the key material, instead of leaking it behind an _owns_secret flag that was never set. The flag and the store-callback "detach" are removed: liboqs ignores a NULL callback, so the detach never took effect and left the native key pointing at a dropped ctypes closure. Related lifecycle fixes: - Public methods raise RuntimeError after free() instead of crashing or, for verify(), silently returning False. - The constructor frees its native allocations if loading the secret key fails, and generate_keypair() releases the key if keygen fails. - Signature.free() and KeyEncapsulation.free() are now idempotent, so calling free() after a with block no longer double frees. Tests track every native allocation and assert each is freed exactly once, and check that each public method raises after free(). Co-Authored-By: Claude Opus 5.5 Signed-off-by: Douglas Stebila --- oqs/oqs.py | 87 ++++++++++++++++++++--------- tests/test_kem.py | 22 ++++++++ tests/test_sig.py | 22 ++++++++ tests/test_stfl_sig.py | 124 +++++++++++++++++++++++++++++++++++++---- 4 files changed, 218 insertions(+), 37 deletions(-) diff --git a/oqs/oqs.py b/oqs/oqs.py index 11a6866..0fbae05 100644 --- a/oqs/oqs.py +++ b/oqs/oqs.py @@ -553,13 +553,16 @@ def decap_secret(self, ciphertext: Union[int, bytes]) -> bytes: raise RuntimeError(msg) def free(self) -> None: - """Releases the native resources.""" + """Releases the native resources. Calling it more than once is safe.""" + if not self._kem: + return if hasattr(self, "secret_key"): native().OQS_MEM_cleanse( ct.byref(self.secret_key), self._kem.contents.length_secret_key, ) native().OQS_KEM_free(self._kem) + self._kem = None def __repr__(self) -> str: return f"Key encapsulation mechanism: {self._kem.contents.method_name.decode()}" @@ -868,13 +871,16 @@ def verify_with_ctx_str( return rv == OQS_SUCCESS def free(self) -> None: - """Releases the native resources.""" + """Releases the native resources. Calling it more than once is safe.""" + if not self._sig: + return if hasattr(self, "secret_key"): native().OQS_MEM_cleanse( ct.byref(self.secret_key), self._sig.contents.length_secret_key, ) native().OQS_SIG_free(self._sig) + self._sig = None def __repr__(self) -> str: return f"Signature mechanism: {self._sig.contents.method_name.decode()}" @@ -1053,23 +1059,27 @@ def __init__(self, alg_name: str, secret_key: Optional[bytes] = None) -> None: ) raise RuntimeError(msg) + # The instance always owns both native handles; free() releases them. + self._secret_key: ct.c_void_p | None = None + self._used_keys: list[bytes] = [] + self._store_cb: Optional[ct.CFUNCTYPE] = None + self._sig = native().OQS_SIG_STFL_new(ct.create_string_buffer(alg_name.encode())) if not self._sig: msg = f"Could not allocate OQS_SIG_STFL for {alg_name}" raise RuntimeError(msg) - for field, _ctype in self._fields_: - if field == "oid" or field.endswith("cb"): - continue - setattr(self, field, getattr(self._sig.contents, field)) - - self._secret_key: ct.c_void_p | None = None - self._owns_secret = False - self._used_keys: list[bytes] = [] - self._store_cb: Optional[ct.CFUNCTYPE] = None + try: + for field, _ctype in self._fields_: + if field == "oid" or field.endswith("cb"): + continue + setattr(self, field, getattr(self._sig.contents, field)) - if secret_key is not None: - self._load_secret_key(secret_key) + if secret_key is not None: + self._load_secret_key(secret_key) + except BaseException: + self.free() + raise self.details = { "name": self.method_name.decode(), @@ -1081,6 +1091,12 @@ def __init__(self, alg_name: str, secret_key: Optional[bytes] = None) -> None: "length_signature": int(self.length_signature), } + def _check_not_freed(self) -> None: + """Raise if free() has already released the native resources.""" + if not self._sig: + msg = "StatefulSignature has been freed" + raise RuntimeError(msg) + def _attach_store_cb(self) -> None: """Attach a callback to store used keys in the stateful signature.""" @@ -1094,10 +1110,11 @@ def _cb(buf: bytes, length: int, _: ct.c_void_p) -> int: def _new_secret_key(self) -> None: """Create a new secret key for the stateful signature.""" - self._secret_key = native().OQS_SIG_STFL_SECRET_KEY_new(self.method_name) - if not self._secret_key: + secret_key = native().OQS_SIG_STFL_SECRET_KEY_new(self.method_name) + if not secret_key: msg = "Could not allocate OQS_SIG_STFL_SECRET_KEY" raise MemoryError(msg) + self._secret_key = secret_key self._attach_store_cb() def _load_secret_key(self, data: bytes) -> None: @@ -1121,11 +1138,12 @@ def generate_keypair(self) -> bytes: Generate a new keypair for the stateful signature. :raise ValueError: If the keypair has already been generated. - :raise RuntimeError: If the keypair generation fails or if a keypair already exists. + :raise RuntimeError: If the keypair generation fails or the instance has been freed. :return: The generated public key as bytes. """ + self._check_not_freed() if self._secret_key is not None: - msg = "Keypair already generated, call free() to release the secret key" + msg = "Keypair already generated" raise ValueError(msg) sig_struct = self._sig.contents @@ -1138,16 +1156,20 @@ def generate_keypair(self) -> bytes: ) raise RuntimeError(msg) - self._secret_key = native().OQS_SIG_STFL_SECRET_KEY_new(self.method_name) - if not self._secret_key: + secret_key = native().OQS_SIG_STFL_SECRET_KEY_new(self.method_name) + if not secret_key: msg = "Could not allocate OQS_SIG_STFL_SECRET_KEY" raise RuntimeError(msg) + self._secret_key = secret_key self._attach_store_cb() pk_buf = ct.create_string_buffer(sig_struct.length_public_key) rc = native().OQS_SIG_STFL_keypair(self._sig, pk_buf, self._secret_key) if rc != OQS_SUCCESS: + native().OQS_SIG_STFL_SECRET_KEY_free(self._secret_key) + self._secret_key = None + self._store_cb = None msg = "Keypair generation failed" raise RuntimeError(msg) return pk_buf.raw @@ -1158,10 +1180,11 @@ def sign(self, message: bytes) -> bytes: :param message: The message to sign. :raises NotImplementedError: If the method is LMS-based, as it is verify-only supported. - :raises RuntimeError: If the secret key is not initialized. + :raises RuntimeError: If the secret key is not initialized or the instance has been freed. :raises ValueError: If the signing fails. :return: The signature on the message as bytes. """ + self._check_not_freed() if self.method_name.startswith(b"LMS"): msg = "LMS algorithms are verify‑only supported." raise NotImplementedError(msg) @@ -1199,8 +1222,10 @@ def verify(self, message: bytes, signature: bytes, public_key: bytes) -> bool: :param message: The signed message. :param signature: The signature on the message. :param public_key: The signer's public key. + :raises RuntimeError: If the instance has been freed. :return: `True` if the signature is valid, `False` otherwise. """ + self._check_not_freed() msg = ct.create_string_buffer(message, len(message)) sig = ct.create_string_buffer(signature, len(signature)) pk = ct.create_string_buffer(public_key, len(public_key)) @@ -1213,7 +1238,9 @@ def export_secret_key(self) -> bytes: :return: The serialized secret key as bytes. :raises ValueError: If the secret key is not initialized. + :raises RuntimeError: If the instance has been freed. """ + self._check_not_freed() if self._secret_key is None: msg = "Secret key not initialised – call generate_keypair() first" raise ValueError(msg) @@ -1231,6 +1258,7 @@ def export_secret_key(self) -> bytes: def sigs_total(self) -> int: """Get the total number of signatures that can be made with the secret key.""" + self._check_not_freed() total = ct.c_uint64() rc = native().OQS_SIG_STFL_sigs_total(self._sig, ct.byref(total), self._secret_key) if rc != OQS_SUCCESS: @@ -1240,6 +1268,7 @@ def sigs_total(self) -> int: def sigs_remaining(self) -> int: """Get the number of remaining signatures that can be made with the secret key.""" + self._check_not_freed() if self._secret_key is None: msg = "Secret key not initialised – call generate_keypair() first" raise ValueError(msg) @@ -1268,12 +1297,20 @@ def __exit__( self.free() def free(self) -> None: - """Free the native resources.""" - if self._store_cb and self._secret_key: - native().OQS_SIG_STFL_SECRET_KEY_SET_store_cb(self._secret_key, None, None) - self._store_cb = None - if self._secret_key and self._owns_secret: + """ + Free the native resources. Calling it more than once is safe. + + The native secret key is securely erased. Serialized secret keys that + were already returned to Python, by export_secret_key() or through + export_used_keys(), are immutable bytes and are not erased. Export the + secret key before calling free() if it is still needed. + """ + if self._secret_key is not None: native().OQS_SIG_STFL_SECRET_KEY_free(self._secret_key) + self._secret_key = None + # The store callback lives inside the native secret key, so the Python + # closure can only be dropped once that key has been freed. + self._store_cb = None if self._sig: native().OQS_SIG_STFL_free(self._sig) self._sig = None diff --git a/tests/test_kem.py b/tests/test_kem.py index a9b9c58..fcc1ec1 100644 --- a/tests/test_kem.py +++ b/tests/test_kem.py @@ -1,6 +1,7 @@ import os import platform # to learn the OS we're on import random +from unittest import mock import oqs @@ -131,6 +132,27 @@ def test_python_attributes() -> None: raise AssertionError(msg) +def test_free_twice() -> None: + lib = oqs.native() + real_free = lib.OQS_KEM_free + freed: list[object] = [] + + def recording_free(ptr: object) -> None: + freed.append(ptr) + real_free(ptr) + + for alg_name in oqs.get_enabled_kem_mechanisms(): + freed.clear() + with mock.patch.object(lib, "OQS_KEM_free", recording_free): + # Passing a secret key exercises the cleanse path in free(). + with oqs.KeyEncapsulation(alg_name, secret_key=b"\x01") as obj: + pass + obj.free() # Explicit free after the context manager must be a no-op. + if len(freed) != 1: + msg = f"{alg_name}: OQS_KEM_free called {len(freed)} times" + raise AssertionError(msg) + + if __name__ == "__main__": try: import nose2 diff --git a/tests/test_sig.py b/tests/test_sig.py index 51e1b7f..0caa124 100644 --- a/tests/test_sig.py +++ b/tests/test_sig.py @@ -1,5 +1,6 @@ import platform # to learn the OS we're on import random +from unittest import mock import oqs from oqs.oqs import Signature, native @@ -189,6 +190,27 @@ def test_python_attributes() -> None: raise AssertionError(msg) +def test_free_twice() -> None: + lib = oqs.native() + real_free = lib.OQS_SIG_free + freed: list[object] = [] + + def recording_free(ptr: object) -> None: + freed.append(ptr) + real_free(ptr) + + for alg_name in oqs.get_enabled_sig_mechanisms(): + freed.clear() + with mock.patch.object(lib, "OQS_SIG_free", recording_free): + # Passing a secret key exercises the cleanse path in free(). + with oqs.Signature(alg_name, secret_key=b"\x01") as obj: + pass + obj.free() # Explicit free after the context manager must be a no-op. + if len(freed) != 1: + msg = f"{alg_name}: OQS_SIG_free called {len(freed)} times" + raise AssertionError(msg) + + if __name__ == "__main__": try: import nose2 diff --git a/tests/test_stfl_sig.py b/tests/test_stfl_sig.py index 793e039..9dd1c7f 100644 --- a/tests/test_stfl_sig.py +++ b/tests/test_stfl_sig.py @@ -1,6 +1,9 @@ +import contextlib +import ctypes as ct import logging import platform # to learn the OS we're on import random +from collections.abc import Callable, Iterator from pathlib import Path from unittest import mock @@ -159,22 +162,119 @@ def test_python_attributes() -> None: raise AssertionError(msg) -def test_free() -> None: +# Fastest XMSS parameter set, used where a test needs a real secret key. +_KEYGEN_ALG = "XMSS-SHA2_10_256" + +# Each native allocator paired with the function that releases what it returns. +_NATIVE_ALLOCATORS = { + "OQS_SIG_STFL_new": "OQS_SIG_STFL_free", + "OQS_SIG_STFL_SECRET_KEY_new": "OQS_SIG_STFL_SECRET_KEY_free", +} + + +@contextlib.contextmanager +def _track_native_allocations() -> Iterator[dict[str, list[int]]]: + """Record the address of every native struct and secret key allocated or freed.""" lib = oqs.native() - real_free = lib.OQS_SIG_STFL_free - freed: list[object] = [] + log: dict[str, list[int]] = {name: [] for pair in _NATIVE_ALLOCATORS.items() for name in pair} - def recording_free(sig_ptr: object) -> None: - freed.append(sig_ptr) - real_free(sig_ptr) + def wrap(name: str, *, record_result: bool) -> object: + real = getattr(lib, name) - sig = oqs.StatefulSignature(oqs.get_enabled_stateful_sig_mechanisms()[0]) - with mock.patch.object(lib, "OQS_SIG_STFL_free", recording_free): - sig.free() - # The struct must not be freed a second time. - sig.free() + def wrapper(*args: object) -> object: + result = real(*args) + ptr = result if record_result else args[0] + address = ct.cast(ptr, ct.c_void_p).value + if address: + log[name].append(address) + return result + + return wrapper + + with contextlib.ExitStack() as stack: + for new, free in _NATIVE_ALLOCATORS.items(): + stack.enter_context(mock.patch.object(lib, new, wrap(new, record_result=True))) + stack.enter_context(mock.patch.object(lib, free, wrap(free, record_result=False))) + yield log + + +def _assert_each_allocation_freed_once(log: dict[str, list[int]]) -> None: + for new, free in _NATIVE_ALLOCATORS.items(): + if sorted(log[new]) != sorted(log[free]): + msg = f"{new} returned {log[new]} but {free} received {log[free]}" + raise AssertionError(msg) + + +def test_free() -> None: + for alg_name in oqs.get_enabled_stateful_sig_mechanisms(): + with _track_native_allocations() as log: + sig = oqs.StatefulSignature(alg_name) + sig.free() + sig.free() # A second call must not free anything again. + if not log["OQS_SIG_STFL_new"]: + msg = f"No OQS_SIG_STFL struct was allocated for {alg_name}" + raise AssertionError(msg) + _assert_each_allocation_freed_once(log) + + +def test_free_with_keypair() -> None: + if _KEYGEN_ALG not in oqs.get_enabled_stateful_sig_mechanisms(): + return + if any(item in _KEYGEN_ALG for item in disabled_sig_patterns): + return + with _track_native_allocations() as log: + with oqs.StatefulSignature(_KEYGEN_ALG) as sig: + sig.generate_keypair() + sig.free() # Explicit free after the context manager must be a no-op. + if not log["OQS_SIG_STFL_SECRET_KEY_new"]: + msg = "No secret key was allocated" + raise AssertionError(msg) + _assert_each_allocation_freed_once(log) - assert len(freed) == 1 # noqa: S101 + +def test_constructor_failure_frees() -> None: + for alg_name in oqs.get_enabled_stateful_sig_mechanisms(): + with _track_native_allocations() as log: + try: + oqs.StatefulSignature(alg_name, secret_key=b"not a secret key") + except ValueError: + pass + else: + msg = f"An invalid secret key was accepted for {alg_name}" + raise AssertionError(msg) + if not log["OQS_SIG_STFL_SECRET_KEY_new"]: + msg = f"No secret key was allocated for {alg_name}" + raise AssertionError(msg) + _assert_each_allocation_freed_once(log) + + +def _assert_raises_freed(label: str, call: Callable[[], object]) -> None: + try: + call() + except RuntimeError as ex: + if "freed" not in str(ex): + msg = f"{label}() raised an unexpected error: {ex}" + raise AssertionError(msg) from ex + else: + msg = f"{label}() did not raise after free()" + raise AssertionError(msg) + + +def test_use_after_free() -> None: + message, signature, public_key = b"message", b"signature", b"public key" + for alg_name in oqs.get_enabled_stateful_sig_mechanisms(): + sig = oqs.StatefulSignature(alg_name) + sig.free() + calls: dict[str, Callable[[], object]] = { + "generate_keypair": sig.generate_keypair, + "sign": lambda s=sig: s.sign(message), + "verify": lambda s=sig: s.verify(message, signature, public_key), + "export_secret_key": sig.export_secret_key, + "sigs_total": sig.sigs_total, + "sigs_remaining": sig.sigs_remaining, + } + for name, call in calls.items(): + _assert_raises_freed(f"{alg_name}.{name}", call) if __name__ == "__main__":