tally
shamir's secret sharing over gf(256).
git clone https://owenrusk.dev/tally.git
commit 7f8992f16014e41668a581b9dcf782d314797c34 parent 5b8432e081a883435eb4029a6b16bb91e6c18c35 author Owen Rusk <owen@papermothgames.com> date 2025-06-11 20:40:51 -0500
tests: the cli, end to end through a subprocess
| tests/test_cli.py | +70 | -0 |
diff --git a/tests/test_cli.py b/tests/test_cli.py
new file mode 100644
index 0000000..0e852cc
--- /dev/null
+++ b/tests/test_cli.py
@@ -0,0 +1,70 @@
+import os
+import subprocess
+import sys
+import tempfile
+import unittest
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parent.parent
+
+
+def tally(*args: str, stdin: bytes = b"") -> subprocess.CompletedProcess[bytes]:
+ return subprocess.run(
+ [sys.executable, "-m", "tally", *args],
+ input=stdin,
+ capture_output=True,
+ cwd=ROOT,
+ )
+
+
+class CliTest(unittest.TestCase):
+ def test_stdin_to_arguments(self) -> None:
+ split = tally("split", "-k", "3", "-n", "5", stdin=b"not a real secret")
+ self.assertEqual(split.returncode, 0, split.stderr)
+ lines = split.stdout.decode().split()
+ self.assertEqual(len(lines), 5)
+ join = tally("join", *lines[1:4])
+ self.assertEqual(join.returncode, 0, join.stderr)
+ self.assertEqual(join.stdout, b"not a real secret")
+
+ def test_file_to_stdin(self) -> None:
+ secret = bytes(range(256)) + os.urandom(100)
+ with tempfile.TemporaryDirectory() as tmp:
+ path = Path(tmp, "secret")
+ path.write_bytes(secret)
+ split = tally("split", "-k", "2", "-n", "3", str(path))
+ self.assertEqual(split.returncode, 0, split.stderr)
+ join = tally("join", stdin=split.stdout)
+ self.assertEqual(join.returncode, 0, join.stderr)
+ self.assertEqual(join.stdout, secret)
+
+ def test_secret_option(self) -> None:
+ split = tally("split", "-k", "2", "-n", "2", "--secret", "not a real secret")
+ join = tally("join", stdin=split.stdout)
+ self.assertEqual(join.stdout, b"not a real secret")
+
+ def test_fewer_than_k(self) -> None:
+ split = tally("split", "-k", "3", "-n", "5", stdin=b"not a real secret")
+ join = tally("join", *split.stdout.decode().split()[:2])
+ self.assertEqual(join.returncode, 1)
+ self.assertEqual(join.stdout, b"")
+ self.assertEqual(join.stderr.decode(), "tally: need 3 shares, got 2\n")
+
+ def test_different_splits(self) -> None:
+ a = tally("split", "-k", "2", "-n", "2", stdin=b"not a real secret").stdout.decode().split()
+ b = tally("split", "-k", "2", "-n", "2", stdin=b"not a real secret").stdout.decode().split()
+ join = tally("join", a[0], b[1])
+ self.assertEqual(join.returncode, 1)
+ self.assertIn("different splits", join.stderr.decode())
+
+ def test_plain_errors(self) -> None:
+ for args in [("split", "-k", "4", "-n", "3"), ("split", "-k", "2", "-n", "3", "/nonexistent"), ("join", "hello")]:
+ result = tally(*args, stdin=b"x")
+ with self.subTest(args=args):
+ self.assertEqual(result.returncode, 1)
+ self.assertTrue(result.stderr.startswith(b"tally: "), result.stderr)
+ self.assertNotIn(b"Traceback", result.stderr)
+
+
+if __name__ == "__main__":
+ unittest.main()