owenrusk.dev

tally

shamir's secret sharing over gf(256).

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

commit b70068ec35ede53ee74700568ae2b1972e29956c
parent 0370c7202a2935356921e3e96c8fc513ee1521ff
author Owen Rusk <owen@papermothgames.com>
date   2024-06-24 20:30:55 -0500
codec: a share as one line of base32

crockford's alphabet in groups of four, so it can be copied by hand. carries a format version, a random id for the split, k and the share's x.
tally/codec.py+80-0
tests/test_codec.py+43-0
diff --git a/tally/codec.py b/tally/codec.py
new file mode 100644
index 0000000..a6e897f
--- /dev/null
+++ b/tally/codec.py
@@ -0,0 +1,80 @@
+import secrets
+from dataclasses import dataclass
+
+VERSION = 1
+PREFIX = f"t{VERSION}"
+# crockford's base32: digits and lowercase letters, without i, l, o and u
+ALPHABET = "0123456789abcdefghjkmnpqrstvwxyz"
+GROUP = 4
+ID_BYTES = 4
+
+_VALUES = {c: i for i, c in enumerate(ALPHABET)}
+
+
+class ShareError(ValueError):
+    pass
+
+
+@dataclass(frozen=True)
+class Share:
+    id: bytes
+    k: int
+    x: int
+    data: bytes
+
+
+def new_id() -> bytes:
+    return secrets.token_bytes(ID_BYTES)
+
+
+def _b32encode(data: bytes) -> str:
+    out = []
+    acc = bits = 0
+    for byte in data:
+        acc = (acc << 8) | byte
+        bits += 8
+        while bits >= 5:
+            bits -= 5
+            out.append(ALPHABET[(acc >> bits) & 31])
+        acc &= (1 << bits) - 1
+    if bits:
+        out.append(ALPHABET[(acc << (5 - bits)) & 31])
+    return "".join(out)
+
+
+def _b32decode(text: str) -> bytes:
+    out = bytearray()
+    acc = bits = 0
+    for ch in text:
+        value = _VALUES.get(ch)
+        if value is None:
+            raise ShareError(f"{ch!r} can't appear in a share")
+        acc = (acc << 5) | value
+        bits += 5
+        if bits >= 8:
+            bits -= 8
+            out.append((acc >> bits) & 0xFF)
+        acc &= (1 << bits) - 1
+    return bytes(out)
+
+
+def encode(share: Share) -> str:
+    body = _b32encode(share.id + bytes([share.k, share.x]) + share.data)
+    groups = [body[i : i + GROUP] for i in range(0, len(body), GROUP)]
+    return "-".join([PREFIX, *groups])
+
+
+def decode(text: str) -> Share:
+    head, sep, rest = text.strip().lower().partition("-")
+    head = head.strip()
+    if not sep or head[:1] != "t" or not head[1:].isdigit():
+        raise ShareError("doesn't look like a 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:
+        raise ShareError("share is too short")
+    k, x = body[ID_BYTES], body[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 :])
diff --git a/tests/test_codec.py b/tests/test_codec.py
new file mode 100644
index 0000000..619c1a7
--- /dev/null
+++ b/tests/test_codec.py
@@ -0,0 +1,43 @@
+import os
+import unittest
+
+from tally import codec
+from tally.codec import Share, ShareError
+
+SHARE = Share(id=bytes.fromhex("a1b2c3d4"), k=3, x=2, data=b"not a real secret")
+
+
+class CodecTest(unittest.TestCase):
+    def test_round_trip(self) -> None:
+        self.assertEqual(codec.decode(codec.encode(SHARE)), SHARE)
+
+    def test_round_trip_lengths(self) -> None:
+        for length in range(1, 40):
+            share = Share(codec.new_id(), 2, 255, os.urandom(length))
+            self.assertEqual(codec.decode(codec.encode(share)), share)
+
+    def test_looks_right(self) -> None:
+        text = codec.encode(SHARE)
+        self.assertTrue(text.startswith("t1-"))
+        groups = text.split("-")[1:]
+        self.assertTrue(all(len(g) == codec.GROUP for g in groups[:-1]))
+        self.assertTrue(set("".join(groups)) <= set(codec.ALPHABET))
+
+    def test_case_and_spacing_dont_matter(self) -> None:
+        text = codec.encode(SHARE)
+        self.assertEqual(codec.decode(text.upper()), SHARE)
+        self.assertEqual(codec.decode("  " + text.replace("-", " - ") + "\n"), SHARE)
+
+    def test_not_a_share(self) -> None:
+        for text in ["", "hello", "t1", "x1-abcd", "t1-ab", "t1-abcd-u000"]:
+            with self.subTest(text=text), self.assertRaises(ShareError):
+                codec.decode(text)
+
+    def test_other_version(self) -> None:
+        text = "t2" + codec.encode(SHARE)[2:]
+        with self.assertRaisesRegex(ShareError, "format 2"):
+            codec.decode(text)
+
+
+if __name__ == "__main__":
+    unittest.main()