owenrusk.dev

tally

shamir's secret sharing over gf(256).

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

tally / tests/test_cli.py -rw-r--r-- · 4507 bytes

  1 import os
  2 import subprocess
  3 import sys
  4 import tempfile
  5 import unittest
  6 from pathlib import Path
  7 
  8 ROOT = Path(__file__).resolve().parent.parent
  9 
 10 
 11 def tally(*args: str, stdin: bytes = b"") -> subprocess.CompletedProcess[bytes]:
 12     return subprocess.run(
 13         [sys.executable, "-m", "tally", *args],
 14         input=stdin,
 15         capture_output=True,
 16         cwd=ROOT,
 17     )
 18 
 19 
 20 class CliTest(unittest.TestCase):
 21     def test_stdin_to_arguments(self) -> None:
 22         split = tally("split", "-k", "3", "-n", "5", stdin=b"not a real secret")
 23         self.assertEqual(split.returncode, 0, split.stderr)
 24         lines = split.stdout.decode().split()
 25         self.assertEqual(len(lines), 5)
 26         join = tally("join", *lines[1:4])
 27         self.assertEqual(join.returncode, 0, join.stderr)
 28         self.assertEqual(join.stdout, b"not a real secret")
 29 
 30     def test_file_to_stdin(self) -> None:
 31         secret = bytes(range(256)) + os.urandom(100)
 32         with tempfile.TemporaryDirectory() as tmp:
 33             path = Path(tmp, "secret")
 34             path.write_bytes(secret)
 35             split = tally("split", "-k", "2", "-n", "3", str(path))
 36         self.assertEqual(split.returncode, 0, split.stderr)
 37         join = tally("join", stdin=split.stdout)
 38         self.assertEqual(join.returncode, 0, join.stderr)
 39         self.assertEqual(join.stdout, secret)
 40 
 41     def test_stdin_skips_blanks_and_comments(self) -> None:
 42         split = tally("split", "-k", "2", "-n", "3", stdin=b"not a real secret")
 43         a, _, c = split.stdout.decode().split()
 44         notes = f"# first one\n{a}\n\n   # the third\n  {c}  \n\n".encode()
 45         join = tally("join", stdin=notes)
 46         self.assertEqual(join.returncode, 0, join.stderr)
 47         self.assertEqual(join.stdout, b"not a real secret")
 48 
 49     def test_no_secret_on_the_command_line(self) -> None:
 50         split = tally("split", "-k", "2", "-n", "2", "--secret", "not a real secret")
 51         self.assertEqual(split.returncode, 2)
 52         self.assertEqual(split.stdout, b"")
 53 
 54     def test_fewer_than_k(self) -> None:
 55         split = tally("split", "-k", "3", "-n", "5", stdin=b"not a real secret")
 56         join = tally("join", *split.stdout.decode().split()[:2])
 57         self.assertEqual(join.returncode, 1)
 58         self.assertEqual(join.stdout, b"")
 59         self.assertEqual(join.stderr.decode(), "tally: need 3 shares, got 2\n")
 60 
 61     def test_different_splits(self) -> None:
 62         a = tally("split", "-k", "2", "-n", "2", stdin=b"not a real secret").stdout.decode().split()
 63         b = tally("split", "-k", "2", "-n", "2", stdin=b"not a real secret").stdout.decode().split()
 64         join = tally("join", a[0], b[1])
 65         self.assertEqual(join.returncode, 1)
 66         self.assertIn("different splits", join.stderr.decode())
 67 
 68     def test_out_dir(self) -> None:
 69         with tempfile.TemporaryDirectory() as tmp:
 70             out = Path(tmp, "shares")
 71             split = tally("split", "-k", "2", "-n", "3", "-o", str(out), stdin=b"not a real secret")
 72             self.assertEqual(split.returncode, 0, split.stderr)
 73             self.assertEqual(split.stdout, b"")
 74             files = sorted(out.iterdir())
 75             self.assertEqual([f.name for f in files], ["share-1.txt", "share-2.txt", "share-3.txt"])
 76             for f in files:
 77                 self.assertEqual(f.stat().st_mode & 0o777, 0o600)
 78             join = tally("join", stdin=files[0].read_bytes() + files[2].read_bytes())
 79             self.assertEqual(join.stdout, b"not a real secret")
 80 
 81     def test_out_dir_wont_overwrite(self) -> None:
 82         with tempfile.TemporaryDirectory() as tmp:
 83             Path(tmp, "share-2.txt").write_text("keep me\n")
 84             split = tally("split", "-k", "2", "-n", "3", "-o", tmp, stdin=b"not a real secret")
 85             self.assertEqual(split.returncode, 1)
 86             self.assertIn(b"not overwriting", split.stderr)
 87             self.assertEqual(os.listdir(tmp), ["share-2.txt"])
 88             self.assertEqual(Path(tmp, "share-2.txt").read_text(), "keep me\n")
 89 
 90     def test_plain_errors(self) -> None:
 91         for args in [("split", "-k", "4", "-n", "3"), ("split", "-k", "2", "-n", "3", "/nonexistent"), ("join", "hello")]:
 92             result = tally(*args, stdin=b"x")
 93             with self.subTest(args=args):
 94                 self.assertEqual(result.returncode, 1)
 95                 self.assertTrue(result.stderr.startswith(b"tally: "), result.stderr)
 96                 self.assertNotIn(b"Traceback", result.stderr)
 97 
 98 
 99 if __name__ == "__main__":
100     unittest.main()