diff --git a/multihash/funcs.py b/multihash/funcs.py index 2304eb2..7ff9e96 100644 --- a/multihash/funcs.py +++ b/multihash/funcs.py @@ -19,9 +19,10 @@ import hashlib from collections import namedtuple +from collections.abc import Callable from enum import IntEnum from numbers import Integral -from typing import ClassVar +from typing import Any, ClassVar import blake3 import mmh3 @@ -228,6 +229,7 @@ def __init__(self, shake_func, length: int): self._shake_func = shake_func self._hasher = shake_func() self._length = length + self.digest_size = length self.name = self._hasher.name def update(self, data: bytes) -> None: @@ -543,6 +545,8 @@ class FuncReg(metaclass=_FuncRegMeta): _func_from_name: ClassVar[dict] = {} _func_from_hash: ClassVar[dict] = {} _func_hash: ClassVar[dict] = {} + default_lengths: ClassVar[dict[int, int]] = {} + _variable_factories: ClassVar[dict[int, Callable]] = {} # Standard hash function data: (func, hashlib_name, constructor) _std_func_data: ClassVar[list] = [ @@ -554,8 +558,7 @@ class FuncReg(metaclass=_FuncRegMeta): (Func.sha3_384, "sha3_384", hashlib.sha3_384), (Func.sha3_256, "sha3_256", hashlib.sha3_256), (Func.sha3_224, "sha3_224", hashlib.sha3_224), - (Func.shake_128, "shake_128", None), # Variable length - use ShakeHash wrapper - (Func.shake_256, "shake_256", None), # Variable length - use ShakeHash wrapper + # SHAKE registered via register_variable_size in reset() (Func.blake2b_256, "blake2b", lambda: hashlib.blake2b(digest_size=32)), (Func.blake2b_512, "blake2b", lambda: hashlib.blake2b(digest_size=64)), (Func.blake2s_256, "blake2s", lambda: hashlib.blake2s(digest_size=32)), @@ -628,6 +631,8 @@ def reset(cls) -> None: cls._func_from_name = {} cls._func_from_hash = {} cls._func_hash = {} + cls.default_lengths = {} + cls._variable_factories = {} for func, hash_name, hash_new in cls._std_func_data: cls._do_register(func, func.name, hash_name, hash_new) @@ -656,6 +661,22 @@ def reset(cls) -> None: # Variant not available, skip pass + # Variable-size SHAKE functions + def shake_128_factory(size_hint: int): + length = 32 if size_hint < 0 else size_hint + if length <= 0: + return None + return ShakeHash(hashlib.shake_128, length) + + def shake_256_factory(size_hint: int): + length = 64 if size_hint < 0 else size_hint + if length <= 0: + return None + return ShakeHash(hashlib.shake_256, length) + + cls.register_variable_size(Func.shake_128, "shake_128", shake_128_factory, "shake_128") + cls.register_variable_size(Func.shake_256, "shake_256", shake_256_factory, "shake_256") + @classmethod def get(cls, func_hint: Func | str | int) -> Func | int: """Return a registered hash function matching the given hint.""" @@ -683,6 +704,44 @@ def _do_register(cls, code: int, name: str, hash_name: str | None = None, hash_n if hash_name: cls._func_from_hash[hash_name] = code cls._func_hash[code] = cls._hash(hash_name, hash_new) + if hash_new is not None and int(code) not in cls._variable_factories: + try: + cls.default_lengths[int(code)] = hash_new().digest_size + except (AttributeError, TypeError, ValueError): + pass + + @classmethod + def register_variable_size( + cls, + code: int, + name: str, + factory: Callable[[int], Any], + hash_name: str | None = None, + ) -> None: + """Register a variable-size hash function factory. + + The factory receives a size hint in bytes. Pass ``-1`` to request the + default output length. The factory should return a hashlib-compatible + object, or ``None`` if the requested size is invalid. + + Args: + code: Multihash function code + name: Function name (hyphen or underscore form) + factory: Callable ``(size_hint: int) -> hash_obj | None`` + hash_name: Optional hashlib-style name for reverse lookup + """ + code = int(code) + default_hasher = factory(-1) + if default_hasher is None: + raise ValueError(f"variable-size factory for {name} rejected default size hint -1") + cls._variable_factories[code] = factory + cls.default_lengths[code] = default_hasher.digest_size + cls._do_register(code, name, hash_name or name, None) + + @classmethod + def is_variable_size(cls, func: Func | int) -> bool: + """Return True if ``func`` was registered via :meth:`register_variable_size`.""" + return int(func) in cls._variable_factories @classmethod def register(cls, code: int, name: str, hash_name: str | None = None, hash_new=None) -> None: @@ -735,6 +794,8 @@ def unregister(cls, code: int) -> None: hash_data = cls._func_hash.pop(code) if hash_data.name: del cls._func_from_hash[hash_data.name] + cls.default_lengths.pop(code, None) + cls._variable_factories.pop(code, None) @classmethod def func_from_hash(cls, hash_obj) -> Func | int: @@ -760,27 +821,25 @@ def hash_from_func(cls, func: Func | int, length: int | None = None): Args: func: Hash function code or Func enum - length: Optional length for SHAKE hashes. Required for SHAKE. Returns None if None for SHAKE. + length: Optional output length for variable-size hashes Returns: Hash object or None if not available Note: - SHAKE functions (shake_128, shake_256) require a length parameter - to specify the output digest size. If length is None for SHAKE - functions, this method returns None. + For variable-size functions, ``None`` returns ``None`` (caller must + supply a length or use ``-1`` for the default). ``-1`` requests the + default length from :attr:`default_lengths`. """ + code = int(func) + if code in cls._variable_factories: + if length is None: + return None + size_hint = -1 if length == -1 else length + return cls._variable_factories[code](size_hint) + new = cls._func_hash[func].new if new is None: - # Handle SHAKE functions with variable length - if func == Func.shake_128: - if length is None: - return None - return ShakeHash(hashlib.shake_128, length) - elif func == Func.shake_256: - if length is None: - return None - return ShakeHash(hashlib.shake_256, length) return None return new() diff --git a/multihash/multihash.py b/multihash/multihash.py index 8d04f98..6c0456a 100644 --- a/multihash/multihash.py +++ b/multihash/multihash.py @@ -19,19 +19,18 @@ from . import base58 -def _resolve_shake_length(func: Func | int, length: int | None) -> int: - """Resolve SHAKE hash output length. +def _resolve_variable_length(func: Func | int, length: int | None) -> int: + """Resolve output length for a variable-size hash function. Args: - func: SHAKE function (shake_128 or shake_256) + func: Variable-size hash function code length: Requested length (None or -1 for default) Returns: Resolved length in bytes """ if length is None or length == -1: - # Default length for SHAKE-128 is 32, SHAKE-256 is 64 - return 32 if func == Func.shake_128 else 64 + return FuncReg.default_lengths[int(func)] return length @@ -669,12 +668,12 @@ def _do_digest(data, func, length: int | None = None): TruncationError: If truncation length is invalid """ func = FuncReg.get(func) - is_shake = func in (Func.shake_128, Func.shake_256) + is_variable = FuncReg.is_variable_size(func) - # Handle SHAKE functions which require length - if is_shake: - shake_length = _resolve_shake_length(func, length) - hash_obj = FuncReg.hash_from_func(func, length=shake_length) + # Handle variable-size functions which require a length hint + if is_variable: + resolved_length = _resolve_variable_length(func, length) + hash_obj = FuncReg.hash_from_func(func, length=resolved_length) else: hash_obj = FuncReg.hash_from_func(func) @@ -684,8 +683,8 @@ def _do_digest(data, func, length: int | None = None): hash_obj.update(data) digest_bytes = bytes(hash_obj.digest()) - # Handle truncation (but not for SHAKE, as they already produce the right length) - if not is_shake and length is not None and length != -1: + # Handle truncation (but not for variable-size digests that already match length) + if not is_variable and length is not None and length != -1: if length < 0: raise TruncationError(f"truncation length must be non-negative, got {length}") if length == 0: @@ -1057,12 +1056,12 @@ def sum_stream( raise ValueError(f"chunk_size must be positive, got {chunk_size}") func = FuncReg.get(code) - is_shake = func in (Func.shake_128, Func.shake_256) + is_variable = FuncReg.is_variable_size(func) - # Handle SHAKE functions which require length - if is_shake: - shake_length = _resolve_shake_length(func, length) - hash_obj = FuncReg.hash_from_func(func, length=shake_length) + # Handle variable-size functions which require a length hint + if is_variable: + resolved_length = _resolve_variable_length(func, length) + hash_obj = FuncReg.hash_from_func(func, length=resolved_length) else: hash_obj = FuncReg.hash_from_func(func) @@ -1078,8 +1077,8 @@ def sum_stream( digest_bytes = bytes(hash_obj.digest()) - # Handle truncation (but not for SHAKE, as they already produce the right length) - if not is_shake and length is not None and length != -1: + # Handle truncation (but not for variable-size digests that already match length) + if not is_variable and length is not None and length != -1: if length < 0: raise TruncationError(f"truncation length must be non-negative, got {length}") if length == 0: diff --git a/newsfragments/55.feature.rst b/newsfragments/55.feature.rst new file mode 100644 index 0000000..e2f091c --- /dev/null +++ b/newsfragments/55.feature.rst @@ -0,0 +1 @@ +Add :meth:`FuncReg.register_variable_size` and :attr:`FuncReg.default_lengths` for variable-output hash registration (SHAKE now uses this path). diff --git a/tests/test_funcs.py b/tests/test_funcs.py index b8ed371..715e869 100644 --- a/tests/test_funcs.py +++ b/tests/test_funcs.py @@ -205,3 +205,47 @@ def test_murmur3_deterministic(self): mh1 = digest(data, "murmur3-128") mh2 = digest(data, "murmur3-128") assert mh1.digest == mh2.digest + + +class VariableSizeRegistryTestCase: + """Tests for FuncReg.register_variable_size and default_lengths.""" + + def test_shake_defaults_in_default_lengths(self): + assert FuncReg.default_lengths[Func.shake_128] == 32 + assert FuncReg.default_lengths[Func.shake_256] == 64 + assert FuncReg.is_variable_size(Func.shake_128) is True + assert FuncReg.is_variable_size(Func.sha2_256) is False + + def test_fixed_hash_default_length(self): + assert FuncReg.default_lengths[Func.sha2_256] == 32 + + def test_register_variable_size_custom(self): + def factory(size_hint: int): + length = 16 if size_hint < 0 else size_hint + if length <= 0 or length > 32: + return None + + class H: + name = "test-var" + digest_size = length + + def __init__(self): + self._data = b"" + + def update(self, data): + self._data += data + + def digest(self): + return (self._data + b"\x00" * length)[:length] + + return H() + + code = 0x07 + try: + FuncReg.register_variable_size(code, "test-var", factory, "test-var") + assert FuncReg.default_lengths[code] == 16 + assert FuncReg.is_variable_size(code) is True + h = FuncReg.hash_from_func(code, length=8) + assert h.digest_size == 8 + finally: + FuncReg.unregister(code)