Skip to content

Commit 47482c2

Browse files
authored
Merge pull request #46 from gaoflow/fix-base10-fast-path
Fix base10 dropping leading zeros and corrupting some values on decode
2 parents 4c234d3 + efeee08 commit 47482c2

2 files changed

Lines changed: 30 additions & 5 deletions

File tree

src/codext/base/_base.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -125,8 +125,10 @@ def base_encode(input, charset, errors="strict", exc=BaseEncodeError):
125125
if i > SIZE_LIMIT:
126126
raise InputSizeLimitError("Input exceeded size limit")
127127
return i * charset[0]
128-
if n == 10:
129-
return str(i) if charset == digits else "".join(charset[int(x)] for x in str(i))
128+
# keep this fast-path only for a non-standard 10-character charset ; the digits
129+
# charset uses the generic loop below (leading zeros, bignums, no str/int round-trip)
130+
if n == 10 and charset != digits:
131+
return "".join(charset[int(x)] for x in str(i))
130132
while i > 0:
131133
i, c = divmod(i, n)
132134
r = charset[c] + r
@@ -148,8 +150,8 @@ def base_decode(input, charset, errors="strict", exc=BaseDecodeError):
148150
i, n, dec = 0, len(charset), lambda n: base_encode(n, [chr(x) for x in range(256)], errors, exc)
149151
if n == 1:
150152
return i2s(len(input))
151-
if n == 10:
152-
return i2s(int(input)) if charset == digits else "".join(str(charset.index(c)) for c in input)
153+
if n == 10 and charset != digits:
154+
return "".join(str(charset.index(c)) for c in input)
153155
for k, c in enumerate(input):
154156
try:
155157
i = i * n + charset.index(c)

tests/test_base.py

Lines changed: 24 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -112,7 +112,30 @@ def test_codec_base8(self):
112112
self.assertEqual(codecs.decode(B8, "base8-01234567"), STR)
113113
self.assertRaises(LookupError, codecs.encode, "test", "base8-0123456")
114114
self.assertRaises(LookupError, codecs.encode, "test", "base8-012345678")
115-
115+
116+
def test_codec_base10(self):
117+
B10 = "2361031878030638688519054699098996"
118+
self.assertEqual(codecs.encode(STR, "base10"), B10)
119+
self.assertEqual(codecs.encode(b(STR), "base10"), b(B10))
120+
self.assertEqual(codecs.decode(B10, "base10"), STR)
121+
self.assertEqual(codecs.decode(b(B10), "base10"), b(STR))
122+
for alias in ["int", "integer", "dec", "decimal"]:
123+
self.assertEqual(codecs.encode(STR, alias), B10)
124+
self.assertEqual(codecs.decode(B10, alias), STR)
125+
# leading null bytes must be preserved as leading '0' characters
126+
self.assertEqual(codecs.encode("\x00abc", "base10"), "06382179")
127+
self.assertEqual(codecs.encode("\x00", "base10"), "0")
128+
self.assertEqual(codecs.encode("\x00\x00abc", "base10"), "006382179")
129+
self.assertEqual(codecs.decode("06382179", "base10"), "\x00abc")
130+
self.assertEqual(codecs.decode("00", "base10"), "\x00\x00")
131+
self.assertEqual(codecs.encode(b("\x00abc"), "base10"), b("06382179"))
132+
self.assertEqual(codecs.decode(b("06382179"), "base10"), b("\x00abc"))
133+
# a value whose big integer ends in a 0xe nibble must survive decoding
134+
self.assertEqual(codecs.decode(codecs.encode(b"d\xc6\xfe", "base10"), "base10"), b"d\xc6\xfe")
135+
# large inputs must not hit the int<->str conversion digit limit
136+
for data in [b"\x00\xff\xfe", b"\x00\x00", b"\x2a" * 2048]:
137+
self.assertEqual(codecs.decode(codecs.encode(data, "base10"), "base10"), data)
138+
116139
def test_codec_base16(self):
117140
B16 = "7468697320697320612074657374"
118141
self.assertEqual(codecs.encode(STR, "base16"), B16)

0 commit comments

Comments
 (0)