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()