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