diff --git a/tests/test_link.py b/tests/test_link.py index d0f9185..1cdc6d3 100644 --- a/tests/test_link.py +++ b/tests/test_link.py @@ -123,6 +123,9 @@ class Connecting(unittest.TestCase): fd.known_hosts(device_id).write_text(f"frame-control-{device_id} ssh-ed25519 AAAA\n") def hosts(self, mapping): + # An IPv4 answer is what ssh is pointed at, so localhost arrives as 127.0.0.1. + if "localhost" in mapping: + mapping = dict({"127.0.0.1": mapping["localhost"]}, **mapping) os.environ["FAKESSH_HOSTS"] = json.dumps(mapping) def calls(self): @@ -136,8 +139,15 @@ class Connecting(unittest.TestCase): def test_falls_through_to_the_address_that_is_really_the_headset(self): # Tried in this order: a name that doesn't resolve, a different device, the headset. - d = self.device("nothing.invalid", "127.0.0.1", "localhost") - self.hosts({"127.0.0.1": "wrong", "localhost": "ok"}) + try: # the "different device": the same port on IPv6 loopback + six = socket.socket(socket.AF_INET6) + self.addCleanup(six.close) + six.bind(("::1", self.port)) + six.listen(4) + except OSError: + self.skipTest("no IPv6 loopback") + d = self.device("nothing.invalid", "::1", "localhost") + self.hosts({"::1": "wrong", "localhost": "ok"}) self.link.connect(["start"]) s = self.link.snapshot() self.assertEqual(s["phase"], "connected", s["error"]) @@ -145,12 +155,13 @@ class Connecting(unittest.TestCase): self.assertEqual([st["state"] for st in s["stages"]], ["done"] * 5) rows = {p["host"]: p for p in s["probes"]} self.assertEqual(rows["nothing.invalid"]["state"], "unresolved") - self.assertEqual(rows["127.0.0.1"]["state"], "sshfailed") - self.assertIn("different headset", rows["127.0.0.1"]["detail"]) + self.assertEqual(rows["::1"]["state"], "sshfailed") + self.assertIn("different headset", rows["::1"]["detail"]) # Every ssh command was pointed at the winner, with the host key pinned per device. alias, opts = self.routes[-1] self.assertEqual(alias, "frame-t") - self.assertIn(f"HostName=localhost", opts) + ip = rows["localhost"]["ip"] # an IPv4 answer is used as is; an IPv6 one keeps the name + self.assertIn("HostName=" + (ip if "." in ip else "localhost"), opts) self.assertIn(f"HostKeyAlias=frame-control-{d['id']}", opts) self.assertIn(f"Port={self.port}", opts) master = [c for c in self.calls() if "ControlMaster=yes" in c][-1] @@ -159,7 +170,7 @@ class Connecting(unittest.TestCase): # It learned: localhost works on this network. learned = {a["host"]: a for a in self.reg.get(d["id"])["addresses"]} self.assertEqual(learned["localhost"]["networks"], ["n-test"]) - self.assertEqual(learned["127.0.0.1"]["networks"], []) + self.assertEqual(learned["::1"]["networks"], []) self.assertTrue(self.link.alive()) self.link.close_master() self.assertFalse(any(p.name.startswith("master-") for p in self.dir.iterdir())) @@ -210,6 +221,14 @@ class Connecting(unittest.TestCase): master = [c for c in self.calls() if "ControlMaster=yes" in c][-1] self.assertIn("StrictHostKeyChecking=yes", master) + def test_ssh_goes_to_the_ipv4_address_that_answered(self): + self.device("localhost") + self.hosts({"localhost": "ok"}) + self.link.connect(["start"]) + s = self.link.snapshot() + self.assertEqual((s["phase"], s["via"]["host"], s["via"]["ip"]), ("connected", "localhost", "127.0.0.1")) + self.assertIn("HostName=127.0.0.1", self.routes[-1][1]) + def test_a_bare_alias_lets_ssh_config_decide(self): self.link.override = "frame-bare" self.hosts({"frame-bare": "ok"}) # the stand-in ssh has no config: the alias is the host @@ -356,7 +375,7 @@ class ServerConnection(unittest.TestCase): " User steamos\nHost *\n# <<< steam-frame (frame) <<<\n") env = {**os.environ, "PYTHONDONTWRITEBYTECODE": "1", "FRAME_CONTROL_SSH_DIR": str(ssh_dir), "FRAME_CONTROL_DATA_DIR": str(cls.dir / "data"), "FAKESSH_LOG": str(cls.dir / "calls.jsonl"), - "FAKESSH_DIR": str(cls.dir), "FAKESSH_HOSTS": json.dumps({"localhost": "ok"}), + "FAKESSH_DIR": str(cls.dir), "FAKESSH_HOSTS": json.dumps({"localhost": "ok", "127.0.0.1": "ok"}), "PATH": f"{FAKESSH}{os.pathsep}{os.environ['PATH']}"} env.pop("FRAME_ALIAS", None) cls.log = tempfile.TemporaryFile()