Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 65 additions & 25 deletions oqs/oqs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()}"
Expand Down Expand Up @@ -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()}"
Expand Down Expand Up @@ -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(),
Expand All @@ -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."""

Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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))
Expand All @@ -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)
Expand All @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -1268,12 +1297,23 @@ 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


native().OQS_SIG_STFL_new.restype = ct.POINTER(StatefulSignature)
Expand Down
22 changes: 22 additions & 0 deletions tests/test_kem.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import os
import platform # to learn the OS we're on
import random
from unittest import mock

import oqs

Expand Down Expand Up @@ -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
Expand Down
22 changes: 22 additions & 0 deletions tests/test_sig.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading