owenrusk.dev

tally

shamir's secret sharing over gf(256).

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

commit 4f5d31046118b3d4e52147a8cb8aa4ac1dada770
parent b70068ec35ede53ee74700568ae2b1972e29956c
author Owen Rusk <owen@papermothgames.com>
date   2024-06-25 21:15:02 -0500
codec: crc32 on the end, so a typo gets caught

a typo used to rebuild the wrong secret without saying so.
tally/codec.py+15-4
tests/test_codec.py+6-0
diff --git a/tally/codec.py b/tally/codec.py
index a6e897f..aa9f9de 100644
--- a/tally/codec.py
+++ b/tally/codec.py
@@ -1,4 +1,5 @@
 import secrets
+import zlib
 from dataclasses import dataclass
 
 VERSION = 1
@@ -7,6 +8,7 @@ PREFIX = f"t{VERSION}"
 ALPHABET = "0123456789abcdefghjkmnpqrstvwxyz"
 GROUP = 4
 ID_BYTES = 4
+CHECK_BYTES = 4
 
 _VALUES = {c: i for i, c in enumerate(ALPHABET)}
 
@@ -58,8 +60,14 @@ def _b32decode(text: str) -> bytes:
     return bytes(out)
 
 
+def _check(payload: bytes) -> bytes:
+    # the prefix is checked too, so a typo there can't pass as another version
+    return zlib.crc32(PREFIX.encode() + payload).to_bytes(CHECK_BYTES, "big")
+
+
 def encode(share: Share) -> str:
-    body = _b32encode(share.id + bytes([share.k, share.x]) + share.data)
+    payload = share.id + bytes([share.k, share.x]) + share.data
+    body = _b32encode(payload + _check(payload))
     groups = [body[i : i + GROUP] for i in range(0, len(body), GROUP)]
     return "-".join([PREFIX, *groups])
 
@@ -72,9 +80,12 @@ def decode(text: str) -> Share:
     if head != PREFIX:
         raise ShareError(f"share format {head[1:]} isn't one this version of tally reads")
     body = _b32decode("".join(rest.split()).replace("-", ""))
-    if len(body) < ID_BYTES + 3:
+    if len(body) < ID_BYTES + 3 + CHECK_BYTES:
         raise ShareError("share is too short")
-    k, x = body[ID_BYTES], body[ID_BYTES + 1]
+    payload, check = body[:-CHECK_BYTES], body[-CHECK_BYTES:]
+    if _check(payload) != check:
+        raise ShareError("checksum doesn't match, look for a typo")
+    k, x = payload[ID_BYTES], payload[ID_BYTES + 1]
     if k < 2 or x == 0:
         raise ShareError("share has a bad k or index")
-    return Share(body[:ID_BYTES], k, x, body[ID_BYTES + 2 :])
+    return Share(payload[:ID_BYTES], k, x, payload[ID_BYTES + 2 :])
diff --git a/tests/test_codec.py b/tests/test_codec.py
index 619c1a7..cb9c106 100644
--- a/tests/test_codec.py
+++ b/tests/test_codec.py
@@ -38,6 +38,12 @@ class CodecTest(unittest.TestCase):
         with self.assertRaisesRegex(ShareError, "format 2"):
             codec.decode(text)
 
+    def test_typo(self) -> None:
+        text = codec.encode(SHARE)
+        typo = text[:10] + ("x" if text[10] != "x" else "y") + text[11:]
+        with self.assertRaisesRegex(ShareError, "checksum"):
+            codec.decode(typo)
+
 
 if __name__ == "__main__":
     unittest.main()