owenrusk.dev

tally

shamir's secret sharing over gf(256).

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

commit 2ae1ae698df6da622b1d52bfca60ed531e8a20e2
parent 8e45be10a979e96100aa5d770c05a7916ca9881a
author Owen Rusk <owen@papermothgames.com>
date   2025-02-20 21:00:12 -0600
join: check shares past the first k against them

extra shares used to be ignored. now each one has to land on the same polynomial, which catches a damaged share that still passes its checksum.
tally/__init__.py+7-1
tally/shamir.py+8-4
tests/test_tally.py+15-0
diff --git a/tally/__init__.py b/tally/__init__.py
index 39edeaf..db41a8b 100644
--- a/tally/__init__.py
+++ b/tally/__init__.py
@@ -34,4 +34,10 @@ def join(texts: Iterable[str]) -> bytes:
         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]])
+    points = [(s.x, s.data) for s in shares]
+    base = points[: first.k]
+    # any share past the first k has to sit on the same polynomial
+    for i, (x, data) in enumerate(points[first.k :], first.k + 1):
+        if shamir.interpolate(base, x) != data:
+            raise ShareError(f"share {i} doesn't agree with the first {first.k}, one of them is damaged")
+    return shamir.combine(base)
diff --git a/tally/shamir.py b/tally/shamir.py
index 4fb7bcb..30d44b4 100644
--- a/tally/shamir.py
+++ b/tally/shamir.py
@@ -23,20 +23,24 @@ def split(secret: bytes, k: int, n: int) -> list[tuple[int, bytes]]:
     return [(x, bytes(ys)) for x, ys in enumerate(shares, 1)]
 
 
-def combine(shares: list[tuple[int, bytes]]) -> bytes:
-    xs = [x for x, _ in shares]
+def interpolate(shares: list[tuple[int, bytes]], x: int = 0) -> bytes:
+    # lagrange through the shares, evaluated at x. x = 0 is the secret.
+    xs = [xj for xj, _ in shares]
     if len(set(xs)) != len(xs):
         raise ValueError("two shares with the same x")
     out = bytearray()
     for i in range(len(shares[0][1])):
-        # lagrange at x = 0
         acc = 0
         for j, (xj, ys) in enumerate(shares):
             num = den = 1
             for m, xm in enumerate(xs):
                 if m != j:
-                    num = gf256.mul(num, xm)
+                    num = gf256.mul(num, xm ^ x)
                     den = gf256.mul(den, xm ^ xj)
             acc ^= gf256.mul(ys[i], gf256.div(num, den))
         out.append(acc)
     return bytes(out)
+
+
+def combine(shares: list[tuple[int, bytes]]) -> bytes:
+    return interpolate(shares, 0)
diff --git a/tests/test_tally.py b/tests/test_tally.py
index 9f3b238..b9129a9 100644
--- a/tests/test_tally.py
+++ b/tests/test_tally.py
@@ -46,6 +46,21 @@ class SplitJoinTest(unittest.TestCase):
         with self.assertRaisesRegex(ShareError, "^share 2: "):
             tally.join([lines[0], typo])
 
+    def test_more_than_k(self) -> None:
+        secret = os.urandom(24)
+        lines = tally.split(secret, 3, 7)
+        self.assertEqual(tally.join(lines), secret)
+        self.assertEqual(tally.join(reversed(lines)), secret)
+
+    def test_damaged_extra_share(self) -> None:
+        # a share that's wrong but still passes its own checksum
+        lines = tally.split(b"not a real secret", 3, 5)
+        share = codec.decode(lines[4])
+        flipped = bytes([share.data[0] ^ 1]) + share.data[1:]
+        bad = codec.encode(codec.Share(share.id, share.k, share.x, flipped))
+        with self.assertRaisesRegex(ShareError, "share 4 doesn't agree with the first 3"):
+            tally.join([*lines[:3], bad])
+
     def test_nothing(self) -> None:
         with self.assertRaises(ShareError):
             tally.split(b"", 2, 3)