diff --git a/multibase/converters.py b/multibase/converters.py index 4202658..f4f0347 100644 --- a/multibase/converters.py +++ b/multibase/converters.py @@ -7,8 +7,17 @@ class BaseStringConverter(BaseConverter): def encode(self, bytes): + bytes = ensure_bytes(bytes) + if len(bytes) == 0: + return b"" + + leading_zeros = len(bytes) - len(bytes.lstrip(b"\x00")) + if leading_zeros == len(bytes): + return str(self.digits[0] * leading_zeros).encode("utf-8") + number = int.from_bytes(bytes, byteorder="big", signed=False) - return ensure_bytes(super().encode(number)) + encoded = super().encode(number) + return ensure_bytes(self.digits[0] * leading_zeros + encoded) def bytes_to_int(self, bytes): length = len(bytes) @@ -19,12 +28,20 @@ def bytes_to_int(self, bytes): value += self.digits.index(chr(x)) * base ** (length - (i + 1)) return value - def decode(self, bytes): - decoded_int = self.bytes_to_int(bytes) + def decode(self, data): + bytes_str = data.decode("utf-8") if isinstance(data, bytes) else data + if len(bytes_str) == 0: + return b"" + + leading_zeros = len(bytes_str) - len(bytes_str.lstrip(self.digits[0])) + if leading_zeros == len(bytes_str): + return b"\x00" * leading_zeros + + decoded_int = self.bytes_to_int(data) # See https://docs.python.org/3.5/library/stdtypes.html#int.to_bytes for more about the magical expression # below decoded_data = decoded_int.to_bytes((decoded_int.bit_length() + 7) // 8, byteorder="big") - return decoded_data + return b"\x00" * leading_zeros + decoded_data class Base16StringConverter(BaseStringConverter): @@ -44,12 +61,7 @@ def decode(self, data): data_str = data.decode("utf-8") else: data_str = data - # Convert to match our digits case - if self.uppercase: - data_str = data_str.upper() - else: - data_str = data_str.lower() - return super().decode(data_str.encode("utf-8")) + return bytes.fromhex(data_str) class BaseByteStringConverter: diff --git a/tests/test_roundtrip.py b/tests/test_roundtrip.py new file mode 100644 index 0000000..4d2ddf5 --- /dev/null +++ b/tests/test_roundtrip.py @@ -0,0 +1,45 @@ +import os + +import pytest + +from multibase import ENCODINGS, decode, encode + + +@pytest.mark.parametrize("encoding_info", ENCODINGS, ids=lambda e: e.encoding) +def test_random_data(encoding_info): + """Round-trip random data of various sizes.""" + for size in [1, 2, 7, 16, 32, 64, 137, 256, 1024]: + data = os.urandom(size) + encoded = encode(encoding_info.encoding, data) + decoded = decode(encoded) + assert decoded == data, f"Failed for {encoding_info.encoding} size={size}" + + +@pytest.mark.parametrize("encoding_info", ENCODINGS, ids=lambda e: e.encoding) +def test_leading_zeros(encoding_info): + """Round-trip data with leading zero bytes.""" + for num_zeros in [1, 2, 4, 8, 16]: + data = b"\x00" * num_zeros + b"hello" + encoded = encode(encoding_info.encoding, data) + decoded = decode(encoded) + assert decoded == data, f"Leading zeros lost for {encoding_info.encoding} zeros={num_zeros}" + + +@pytest.mark.parametrize("encoding_info", ENCODINGS, ids=lambda e: e.encoding) +def test_all_zeros(encoding_info): + """Round-trip all-zero data.""" + for size in [1, 4, 16, 32]: + data = b"\x00" * size + encoded = encode(encoding_info.encoding, data) + decoded = decode(encoded) + assert decoded == data + + +@pytest.mark.parametrize("encoding_info", ENCODINGS, ids=lambda e: e.encoding) +def test_all_ones(encoding_info): + """Round-trip all-0xFF data.""" + for size in [1, 4, 16, 32]: + data = b"\xff" * size + encoded = encode(encoding_info.encoding, data) + decoded = decode(encoded) + assert decoded == data