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 = 0acc = 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)