owenrusk.dev

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