tally
shamir's secret sharing over gf(256).
git clone https://owenrusk.dev/tally.git
commit 7422d9b9ec8d29fe3279d27ff4840d20b8c829f7 parent bdc899c85c20d8e61ea10ea5d4065e663de0b4af author Owen Rusk <owen@papermothgames.com> date 2024-10-23 20:35:09 -0500
move the join checks out of the cli split and join live in the package now, so the checks hold for anything that imports it, not just the command.
| tally/__init__.py | +37 | -0 |
| tally/cli.py | +4 | -18 |
| tests/test_tally.py | +57 | -0 |
diff --git a/tally/__init__.py b/tally/__init__.py index e69de29..39edeaf 100644 --- a/tally/__init__.py +++ b/tally/__init__.py @@ -0,0 +1,37 @@ +from collections.abc import Iterable + +from . import codec, shamir +from .codec import ShareError + +__all__ = ["ShareError", "join", "split"] + + +def split(secret: bytes, k: int, n: int) -> list[str]: + if not secret: + raise ShareError("nothing to split") + split_id = codec.new_id() + return [codec.encode(codec.Share(split_id, k, x, data)) for x, data in shamir.split(secret, k, n)] + + +def join(texts: Iterable[str]) -> bytes: + shares = [] + for i, text in enumerate(texts, 1): + try: + shares.append(codec.decode(text)) + except ShareError as e: + raise ShareError(f"share {i}: {e}") from None + if not shares: + raise ShareError("no shares given") + first = shares[0] + if any(s.id != first.id for s in shares): + raise ShareError("these shares come from different splits") + if any(s.k != first.k or len(s.data) != len(first.data) for s in shares): + raise ShareError("these shares disagree about k or length, one of them is damaged") + seen: dict[int, int] = {} + for i, s in enumerate(shares, 1): + if s.x in seen: + raise ShareError(f"shares {seen[s.x]} and {i} are both share {s.x} of the split") + seen[s.x] = i + if len(shares) < first.k: + raise ShareError(f"need {first.k} shares, got {len(shares)}") + return shamir.combine([(s.x, s.data) for s in shares[: first.k]]) diff --git a/tally/cli.py b/tally/cli.py index f5a9c37..9bd1c1e 100644 --- a/tally/cli.py +++ b/tally/cli.py @@ -1,8 +1,7 @@ import argparse import sys-from . import codec, shamir-from .codec import ShareError+from . import ShareError, join, split def cmd_split(args: argparse.Namespace) -> int: @@ -13,27 +12,14 @@ def cmd_split(args: argparse.Namespace) -> int: else: with open(args.file, "rb") as f: secret = f.read()- split_id = codec.new_id()- for x, data in shamir.split(secret, args.k, args.n):- print(codec.encode(codec.Share(split_id, args.k, x, data)))+ for line in split(secret, args.k, args.n): + print(line) return 0 def cmd_join(args: argparse.Namespace) -> int: lines = args.shares or [line for line in sys.stdin if line.strip()]- shares = [codec.decode(line) for line in lines]- if len({s.id for s in shares}) > 1:- raise ShareError("these shares come from different splits")- k = shares[0].k- seen: dict[int, int] = {}- for i, s in enumerate(shares, 1):- if s.x in seen:- raise ShareError(f"shares {seen[s.x]} and {i} are both share {s.x} of the split")- seen[s.x] = i- if len(shares) < k:- raise ShareError(f"need {k} shares, got {len(shares)}")- secret = shamir.combine([(s.x, s.data) for s in shares[:k]])- sys.stdout.buffer.write(secret)+ sys.stdout.buffer.write(join(lines)) return 0 diff --git a/tests/test_tally.py b/tests/test_tally.py new file mode 100644 index 0000000..9f3b238 --- /dev/null +++ b/tests/test_tally.py @@ -0,0 +1,57 @@ +import itertools +import os +import unittest + +import tally +from tally import ShareError, codec + + +class SplitJoinTest(unittest.TestCase): + def test_round_trip(self) -> None: + for length in [1, 2, 7, 31, 100, 1024]: + secret = os.urandom(length) + lines = tally.split(secret, 3, 5) + with self.subTest(length=length): + self.assertEqual(tally.join(lines[2:]), secret) + + def test_every_k_subset(self) -> None: + secret = os.urandom(40) + lines = tally.split(secret, 3, 6) + for subset in itertools.combinations(lines, 3): + self.assertEqual(tally.join(subset), secret) + + def test_one_id_per_split(self) -> None: + lines = tally.split(b"x", 2, 4) + self.assertEqual(len({codec.decode(line).id for line in lines}), 1) + + def test_fewer_than_k(self) -> None: + lines = tally.split(b"not a real secret", 3, 5) + with self.assertRaisesRegex(ShareError, "need 3 shares, got 2"): + tally.join(lines[:2]) + + def test_same_share_twice(self) -> None: + lines = tally.split(b"not a real secret", 2, 3) + with self.assertRaisesRegex(ShareError, "shares 1 and 2 are both share 1"): + tally.join([lines[0], lines[0]]) + + def test_different_splits(self) -> None: + a = tally.split(b"not a real secret", 2, 3) + b = tally.split(b"not a real secret", 2, 3) + with self.assertRaisesRegex(ShareError, "different splits"): + tally.join([a[0], b[1]]) + + def test_typo_names_the_share(self) -> None: + lines = tally.split(b"not a real secret", 2, 3) + typo = lines[1][:-1] + ("0" if lines[1][-1] != "0" else "2") + with self.assertRaisesRegex(ShareError, "^share 2: "): + tally.join([lines[0], typo]) + + def test_nothing(self) -> None: + with self.assertRaises(ShareError): + tally.split(b"", 2, 3) + with self.assertRaises(ShareError): + tally.join([]) + + +if __name__ == "__main__": + unittest.main()