| 1 | #!/usr/bin/env python3 |
| 2 | """Route .studio.test through the rehearsal tunnels on this Mac; stop restores the prior settings.""" |
| 3 | |
| 4 | import argparse |
| 5 | import asyncio |
| 6 | import base64 |
| 7 | import hashlib |
| 8 | import json |
| 9 | import os |
| 10 | from pathlib import Path |
| 11 | import signal |
| 12 | import socket |
| 13 | import ssl |
| 14 | import subprocess |
| 15 | import sys |
| 16 | import time |
| 17 | |
| 18 | STATE = Path("/var/db/snowglobe-test") |
| 19 | RESOLVER = Path("/etc/resolver/studio.test") |
| 20 | HOSTS = Path("/etc/hosts") |
| 21 | KEYCHAIN = "/Library/Keychains/System.keychain" |
| 22 | FORWARDS = ((80, 27080), (443, 27443)) |
| 23 | |
| 24 | |
| 25 | def local_hosts(content): |
| 26 | lines = [] |
| 27 | aliases = [] |
| 28 | for line in content.splitlines(keepends=True): |
| 29 | names, _, comment = line.partition("#") |
| 30 | fields = names.split() |
| 31 | if len(fields) < 2 or not any(name.endswith(".studio.test") for name in fields[1:]): |
| 32 | lines.append(line) |
| 33 | continue |
| 34 | aliases.extend(name for name in fields[1:] if name.endswith(".studio.test")) |
| 35 | remaining = [name for name in fields[1:] if not name.endswith(".studio.test")] |
| 36 | if remaining: |
| 37 | lines.append(" ".join([fields[0], *remaining]) + (" #" + comment.rstrip() if comment else "") + "\n") |
| 38 | elif comment: |
| 39 | lines.append("#" + comment.rstrip() + "\n") |
| 40 | if aliases: |
| 41 | if lines and not lines[-1].endswith("\n"): |
| 42 | lines[-1] += "\n" |
| 43 | lines.append("127.0.0.1 " + " ".join(dict.fromkeys(aliases)) + "\n") |
| 44 | return "".join(lines) |
| 45 | |
| 46 | |
| 47 | class DNSRelay(asyncio.DatagramProtocol): |
| 48 | def __init__(self): |
| 49 | self.requests = set() |
| 50 | |
| 51 | def connection_made(self, transport): |
| 52 | self.transport = transport |
| 53 | |
| 54 | def datagram_received(self, data, address): |
| 55 | if len(data) < 12 or len(self.requests) >= 64: |
| 56 | return |
| 57 | request = asyncio.create_task(asyncio.wait_for(self.forward(data, address), timeout=3)) |
| 58 | self.requests.add(request) |
| 59 | request.add_done_callback(self.finished) |
| 60 | |
| 61 | def finished(self, request): |
| 62 | self.requests.discard(request) |
| 63 | if not request.cancelled(): |
| 64 | request.exception() |
| 65 | |
| 66 | async def forward(self, data, address): |
| 67 | writer = None |
| 68 | try: |
| 69 | reader, writer = await asyncio.open_connection("127.0.0.1", 53153) |
| 70 | writer.write(len(data).to_bytes(2, "big") + data) |
| 71 | await writer.drain() |
| 72 | length = int.from_bytes(await reader.readexactly(2), "big") |
| 73 | response = await reader.readexactly(length) |
| 74 | if len(response) >= 12 and response[:2] == data[:2]: |
| 75 | self.transport.sendto(response, address) |
| 76 | finally: |
| 77 | if writer: |
| 78 | writer.close() |
| 79 | |
| 80 | |
| 81 | async def relay(listeners): |
| 82 | async def connect(reader, writer, upstream): |
| 83 | peer_writer = None |
| 84 | try: |
| 85 | peer_reader, peer_writer = await asyncio.open_connection("127.0.0.1", upstream) |
| 86 | |
| 87 | async def pipe(source, destination): |
| 88 | while data := await source.read(65536): |
| 89 | destination.write(data) |
| 90 | await destination.drain() |
| 91 | if destination.can_write_eof(): |
| 92 | destination.write_eof() |
| 93 | |
| 94 | await asyncio.gather(pipe(reader, peer_writer), pipe(peer_reader, writer)) |
| 95 | except (OSError, asyncio.CancelledError): |
| 96 | pass |
| 97 | finally: |
| 98 | writer.close() |
| 99 | if peer_writer: |
| 100 | peer_writer.close() |
| 101 | |
| 102 | servers = [] |
| 103 | for listener, upstream in listeners: |
| 104 | servers.append(await asyncio.start_server( |
| 105 | lambda reader, writer, port=upstream: connect(reader, writer, port), sock=listener, |
| 106 | )) |
| 107 | transport, _ = await asyncio.get_running_loop().create_datagram_endpoint( |
| 108 | DNSRelay, local_addr=("127.0.0.1", 53153), |
| 109 | ) |
| 110 | try: |
| 111 | await asyncio.gather(*(server.serve_forever() for server in servers)) |
| 112 | finally: |
| 113 | transport.close() |
| 114 | |
| 115 | |
| 116 | def stop(): |
| 117 | record = STATE / "settings.json" |
| 118 | if not record.exists(): |
| 119 | print("Local test routing is already stopped.") |
| 120 | return |
| 121 | settings = json.loads(record.read_text()) |
| 122 | for name, saved in settings["files"].items(): |
| 123 | path = Path(name) |
| 124 | current = hashlib.sha256(path.read_bytes()).hexdigest() if path.exists() else None |
| 125 | before = hashlib.sha256(base64.b64decode(saved["before"])).hexdigest() if saved["before"] is not None else None |
| 126 | if current not in (before, saved["applied"]): |
| 127 | raise SystemExit(f"{path} changed after setup. Restore it using {record}; nothing was overwritten.") |
| 128 | if pid := settings.get("pid"): |
| 129 | command = subprocess.run(["ps", "-p", str(pid), "-o", "command="], capture_output=True, text=True).stdout |
| 130 | if str(Path(__file__).resolve()) + " --relay " in command: |
| 131 | try: |
| 132 | os.kill(pid, signal.SIGTERM) |
| 133 | except ProcessLookupError: |
| 134 | pass |
| 135 | for name, saved in settings["files"].items(): |
| 136 | path = Path(name) |
| 137 | if saved["before"] is None: |
| 138 | path.unlink(missing_ok=True) |
| 139 | else: |
| 140 | path.write_bytes(base64.b64decode(saved["before"])) |
| 141 | if settings.get("addedCA"): |
| 142 | subprocess.run(["security", "remove-trusted-cert", "-d", str(STATE / "ca.crt")], check=True) |
| 143 | subprocess.run(["security", "delete-certificate", "-Z", settings["sha1"], KEYCHAIN], check=True) |
| 144 | subprocess.run(["dscacheutil", "-flushcache"], check=True) |
| 145 | subprocess.run(["killall", "-HUP", "mDNSResponder"], check=True) |
| 146 | record.unlink() |
| 147 | print("Previous DNS, hosts entries, and certificate trust restored.") |
| 148 | |
| 149 | |
| 150 | def start(certificate): |
| 151 | if (STATE / "settings.json").exists(): |
| 152 | raise SystemExit("Local test routing is already configured. Run stop before starting it again.") |
| 153 | uid, gid = int(os.environ["SUDO_UID"]), int(os.environ["SUDO_GID"]) |
| 154 | if uid == 0: |
| 155 | raise SystemExit("Run this with sudo from your normal Mac account.") |
| 156 | for _, port in FORWARDS: |
| 157 | with socket.create_connection(("127.0.0.1", port), timeout=3): |
| 158 | pass |
| 159 | answer = subprocess.check_output(["dig", "+tcp", "@127.0.0.1", "-p", "53153", "globe.studio.test", "+short"]) |
| 160 | if answer.strip() != b"127.0.0.1": |
| 161 | raise SystemExit("The rehearsal DNS tunnel is unavailable. Start its SSH forwards first.") |
| 162 | context = ssl.create_default_context(cafile=str(certificate)) |
| 163 | with socket.create_connection(("127.0.0.1", 27443), timeout=5) as upstream: |
| 164 | with context.wrap_socket(upstream, server_hostname="globe.studio.test"): |
| 165 | pass |
| 166 | for port, _ in FORWARDS: |
| 167 | with socket.socket() as probe: |
| 168 | probe.bind(("127.0.0.1", port)) |
| 169 | with socket.socket(type=socket.SOCK_DGRAM) as probe: |
| 170 | probe.bind(("127.0.0.1", 53153)) |
| 171 | pem = certificate.read_bytes() |
| 172 | der = subprocess.check_output(["openssl", "x509", "-outform", "DER"], input=pem) |
| 173 | existing = subprocess.run(["security", "find-certificate", "-a", "-p", KEYCHAIN], capture_output=True).stdout |
| 174 | sha1 = hashlib.sha1(der).hexdigest().upper() |
| 175 | installed = False |
| 176 | for part in existing.split(b"-----END CERTIFICATE-----")[:-1]: |
| 177 | found = subprocess.check_output(["openssl", "x509", "-outform", "DER"], input=part + b"-----END CERTIFICATE-----\n") |
| 178 | installed = installed or found == der |
| 179 | trusted = subprocess.run(["security", "verify-cert", "-c", str(certificate), "-p", "ssl"], capture_output=True).returncode == 0 |
| 180 | if installed and not trusted: |
| 181 | raise SystemExit("The rehearsal CA has existing custom trust settings. Enable its SSL trust in Keychain Access before setup.") |
| 182 | STATE.mkdir(mode=0o700, exist_ok=True) |
| 183 | os.chmod(STATE, 0o700) |
| 184 | (STATE / "ca.crt").write_bytes(pem) |
| 185 | replacements = {HOSTS: local_hosts(HOSTS.read_text()).encode(), |
| 186 | RESOLVER: b"nameserver 127.0.0.1\nport 53153\n"} |
| 187 | settings = {"files": {}, "addedCA": False, "sha1": sha1} |
| 188 | for path, content in replacements.items(): |
| 189 | settings["files"][str(path)] = { |
| 190 | "before": base64.b64encode(path.read_bytes()).decode() if path.exists() else None, |
| 191 | "applied": hashlib.sha256(content).hexdigest(), |
| 192 | } |
| 193 | record = STATE / "settings.json" |
| 194 | record.write_text(json.dumps(settings, indent=2) + "\n") |
| 195 | os.chmod(record, 0o600) |
| 196 | try: |
| 197 | for path, content in replacements.items(): |
| 198 | path.parent.mkdir(exist_ok=True) |
| 199 | path.write_bytes(content) |
| 200 | os.chmod(path, 0o644) |
| 201 | if not trusted: |
| 202 | subprocess.run(["security", "add-trusted-cert", "-d", "-r", "trustRoot", "-p", "ssl", "-k", KEYCHAIN, str(STATE / "ca.crt")], check=True) |
| 203 | settings["addedCA"] = True |
| 204 | record.write_text(json.dumps(settings, indent=2) + "\n") |
| 205 | with (STATE / "relay.log").open("ab") as log: |
| 206 | process = subprocess.Popen([sys.executable, str(Path(__file__).resolve()), "--relay", str(uid), str(gid)], |
| 207 | stdin=subprocess.DEVNULL, stdout=log, stderr=subprocess.STDOUT, start_new_session=True) |
| 208 | settings["pid"] = process.pid |
| 209 | record.write_text(json.dumps(settings, indent=2) + "\n") |
| 210 | time.sleep(0.3) |
| 211 | if process.poll() is not None: |
| 212 | raise RuntimeError(f"The HTTPS relay did not start. See {STATE / 'relay.log'}.") |
| 213 | subprocess.run(["dscacheutil", "-flushcache"], check=True) |
| 214 | subprocess.run(["killall", "-HUP", "mDNSResponder"], check=True) |
| 215 | subprocess.run(["curl", "--max-time", "8", "--fail", "--silent", "--show-error", |
| 216 | "--output", "/dev/null", "https://globe.studio.test/"], check=True) |
| 217 | except BaseException: |
| 218 | stop() |
| 219 | raise |
| 220 | print("Local .studio.test routing is active. Open https://globe.studio.test/.") |
| 221 | |
| 222 | |
| 223 | if __name__ == "__main__": |
| 224 | if sys.platform != "darwin" or os.geteuid() != 0: |
| 225 | raise SystemExit("This setup needs sudo on the Mac for its resolver, certificate trust, and loopback ports 80/443.") |
| 226 | if len(sys.argv) == 4 and sys.argv[1] == "--relay": |
| 227 | listeners = [] |
| 228 | for port, upstream in FORWARDS: |
| 229 | listener = socket.socket() |
| 230 | listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) |
| 231 | listener.bind(("127.0.0.1", port)) |
| 232 | listener.listen() |
| 233 | listeners.append((listener, upstream)) |
| 234 | os.setgroups([]) |
| 235 | os.setgid(int(sys.argv[3])) |
| 236 | os.setuid(int(sys.argv[2])) |
| 237 | asyncio.run(relay(listeners)) |
| 238 | else: |
| 239 | parser = argparse.ArgumentParser(description=__doc__) |
| 240 | parser.add_argument("action", choices=("start", "stop")) |
| 241 | parser.add_argument("certificate", type=Path, nargs="?") |
| 242 | args = parser.parse_args() |
| 243 | if args.action == "stop": |
| 244 | stop() |
| 245 | elif args.certificate: |
| 246 | start(args.certificate.resolve()) |
| 247 | else: |
| 248 | parser.error("start needs the rehearsal's public CA certificate") |