owenrusk.dev

tally

shamir's secret sharing over gf(256).

git clone https://owenrusk.dev/tally.git

commit 0c268631da02f3502d28d0712938bb4469b5c8a6
parent b5aff2ecb7500636473960b19110116ad7e0da48
author Owen Rusk <owen@papermothgames.com>
date   2024-08-06 19:45:18 -0500
codec: reject padding bits that aren't zero

a typo in the last character could change only the padding and decode to the same bytes, so the check never saw it. the new test tries every one-character typo.
tally/codec.py+3-0
tests/test_codec.py+19-0
diff --git a/tally/codec.py b/tally/codec.py
index 2748170..6514ee2 100644
--- a/tally/codec.py
+++ b/tally/codec.py
@@ -59,6 +59,9 @@ def _b32decode(text: str) -> bytes:
             bits -= 8
             out.append((acc >> bits) & 0xFF)
         acc &= (1 << bits) - 1
+    # what's left over is padding: under five bits, and all zero
+    if bits >= 5 or acc:
+        raise ShareError("share is the wrong length, or its last character is off")
     return bytes(out)
 
 
diff --git a/tests/test_codec.py b/tests/test_codec.py
index 214848b..d97b365 100644
--- a/tests/test_codec.py
+++ b/tests/test_codec.py
@@ -53,6 +53,25 @@ class CodecTest(unittest.TestCase):
         with self.assertRaisesRegex(ShareError, "checksum"):
             codec.decode(typo)
 
+    def test_every_single_character_typo_is_caught(self) -> None:
+        text = codec.encode(SHARE)
+        for i, c in enumerate(text):
+            if c == "-":
+                continue
+            for other in codec.ALPHABET:
+                if other == c:
+                    continue
+                typo = text[:i] + other + text[i + 1 :]
+                with self.subTest(i=i, other=other), self.assertRaises(ShareError):
+                    codec.decode(typo)
+
+    def test_dropped_character_is_caught(self) -> None:
+        text = codec.encode(SHARE)
+        for i, c in enumerate(text):
+            if c != "-":
+                with self.subTest(i=i), self.assertRaises(ShareError):
+                    codec.decode(text[:i] + text[i + 1 :])
+
 
 if __name__ == "__main__":
     unittest.main()