| 1 | #!/usr/bin/env python3 |
| 2 | """Exercise the bootstrap's real SSH preflight and firewall block with a fake UFW.""" |
| 3 | import os |
| 4 | from pathlib import Path |
| 5 | import subprocess |
| 6 | import tempfile |
| 7 | import unittest |
| 8 | |
| 9 | |
| 10 | class BootstrapSshTests(unittest.TestCase): |
| 11 | def run_policy(self, cidrs=None, allow_any=None): |
| 12 | source = Path(__file__).with_name("bootstrap-ubuntu.sh").read_text() |
| 13 | preflight = source.split("# SSH_ALLOWED_CIDRS accepts", 1)[1].split("apt-get update", 1)[0] |
| 14 | firewall = source.split("# Apply the already validated narrow rules", 1)[1].split("ufw --force enable", 1)[0] |
| 15 | with tempfile.TemporaryDirectory() as directory: |
| 16 | receipt = Path(directory) / "ufw.txt" |
| 17 | env = dict(os.environ, UFW_RECEIPT=str(receipt)) |
| 18 | env.pop("SSH_ALLOWED_CIDRS", None) |
| 19 | env.pop("SSH_ALLOW_ANY_SOURCE", None) |
| 20 | if cidrs is not None: |
| 21 | env["SSH_ALLOWED_CIDRS"] = cidrs |
| 22 | if allow_any is not None: |
| 23 | env["SSH_ALLOW_ANY_SOURCE"] = allow_any |
| 24 | # Only the actual preflight and firewall policy execute. Package |
| 25 | # installation, user creation, cloning and host files never run. |
| 26 | script = "set -euo pipefail\nufw() { printf '%s\\n' \"$*\" >>\"$UFW_RECEIPT\"; }\n" |
| 27 | script += "# SSH_ALLOWED_CIDRS accepts" + preflight |
| 28 | script += "# Apply the already validated narrow rules" + firewall |
| 29 | result = subprocess.run(["bash", "-c", script], env=env, capture_output=True, text=True) |
| 30 | return result, receipt.read_text().splitlines() if receipt.exists() else [] |
| 31 | |
| 32 | def test_default_rejects_before_any_firewall_change(self): |
| 33 | result, calls = self.run_policy() |
| 34 | self.assertNotEqual(result.returncode, 0) |
| 35 | self.assertIn("No host changes", result.stderr) |
| 36 | self.assertEqual(calls, []) |
| 37 | |
| 38 | def test_whitespace_and_commas_do_not_opt_in(self): |
| 39 | result, calls = self.run_policy(" ,\t,\n ") |
| 40 | self.assertNotEqual(result.returncode, 0) |
| 41 | self.assertEqual(calls, []) |
| 42 | |
| 43 | def test_explicit_public_ssh_is_warned(self): |
| 44 | result, calls = self.run_policy(allow_any="1") |
| 45 | self.assertEqual(result.returncode, 0, result.stderr) |
| 46 | self.assertIn("reachable from every source", result.stderr) |
| 47 | self.assertEqual(calls, ["allow OpenSSH"]) |
| 48 | |
| 49 | def test_mistyped_public_opt_in_is_denied(self): |
| 50 | result, calls = self.run_policy(allow_any="true") |
| 51 | self.assertNotEqual(result.returncode, 0) |
| 52 | self.assertEqual(calls, []) |
| 53 | |
| 54 | def test_narrow_ipv4_and_ipv6_are_added_before_broad_rule_removal(self): |
| 55 | result, calls = self.run_policy("203.0.113.4/32, 2001:db8::/64") |
| 56 | self.assertEqual(result.returncode, 0, result.stderr) |
| 57 | self.assertEqual(calls, ["allow from 203.0.113.4/32 to any app OpenSSH", "allow from 2001:db8::/64 to any app OpenSSH", "delete allow OpenSSH"]) |
| 58 | |
| 59 | def test_explicit_cidrs_take_precedence_over_public_opt_in(self): |
| 60 | result, calls = self.run_policy("203.0.113.0/24", "1") |
| 61 | self.assertEqual(result.returncode, 0, result.stderr) |
| 62 | self.assertNotIn("allow OpenSSH", calls) |
| 63 | |
| 64 | def test_invalid_later_cidr_changes_no_rules(self): |
| 65 | result, calls = self.run_policy("203.0.113.4/32, ::::::/64") |
| 66 | self.assertNotEqual(result.returncode, 0) |
| 67 | self.assertEqual(calls, []) |
| 68 | |
| 69 | def test_unbounded_or_invalid_networks_are_denied(self): |
| 70 | for cidr in ["0.0.0.0/0", "::/0", "203.0.113.4", "300.1.2.3/32", "2001:db8::/129", "fe80::%en0/64"]: |
| 71 | with self.subTest(cidr=cidr): |
| 72 | result, calls = self.run_policy(cidr) |
| 73 | self.assertNotEqual(result.returncode, 0) |
| 74 | self.assertEqual(calls, []) |
| 75 | |
| 76 | |
| 77 | if __name__ == "__main__": |
| 78 | unittest.main() |
| 79 |