diff --git a/multihash/multihash.py b/multihash/multihash.py index 4472b42..a6ec3b6 100644 --- a/multihash/multihash.py +++ b/multihash/multihash.py @@ -593,6 +593,11 @@ def _do_digest(data, func, length: int | None = None): raise TruncationError(f"truncation length must be non-negative, got {length}") if length == 0: raise TruncationError("truncation length cannot be zero") + if func == Func.identity and length != len(digest_bytes): + raise TruncationError( + f"the length of the identity hash ({length}) must be equal " + f"to the length of the data ({len(digest_bytes)})" + ) if length > len(digest_bytes): raise TruncationError(f"truncation length {length} exceeds digest size {len(digest_bytes)}") digest_bytes = digest_bytes[:length] @@ -935,6 +940,11 @@ def sum_stream( raise TruncationError(f"truncation length must be non-negative, got {length}") if length == 0: raise TruncationError("truncation length cannot be zero") + if func == Func.identity and length != len(digest_bytes): + raise TruncationError( + f"the length of the identity hash ({length}) must be equal " + f"to the length of the data ({len(digest_bytes)})" + ) if length > len(digest_bytes): raise TruncationError(f"truncation length {length} exceeds digest size {len(digest_bytes)}") digest_bytes = digest_bytes[:length] diff --git a/newsfragments/50.bugfix.rst b/newsfragments/50.bugfix.rst new file mode 100644 index 0000000..a3a60fc --- /dev/null +++ b/newsfragments/50.bugfix.rst @@ -0,0 +1 @@ +Reject identity multihashes when a truncation ``length`` is provided that does not equal the input data length, matching go-multihash behavior. diff --git a/tests/test_new_features.py b/tests/test_new_features.py index caba95c..7ed8cfd 100644 --- a/tests/test_new_features.py +++ b/tests/test_new_features.py @@ -151,6 +151,19 @@ def test_verify_with_truncated_digest(self): assert mh.verify(b"hello") is True assert mh.verify(b"world") is False + def test_identity_hash_length_must_match(self): + """Identity hash length must equal data length when length is set.""" + mh = sum(b"hello", Func.identity, length=5) + assert mh.digest == b"hello" + + with pytest.raises(TruncationError, match="identity hash"): + sum(b"hello", Func.identity, length=3) + + def test_identity_hash_length_must_match_sum_stream(self): + """sum_stream rejects mismatched identity truncation length.""" + with pytest.raises(TruncationError, match="identity hash"): + sum_stream(BytesIO(b"hello"), Func.identity, length=3) + class TestShakeFunctionsTestCase: """Tests for SHAKE variable-length hash functions."""