From 13eb65603bc4be5b98db00f41a41a7c111f0e9ed Mon Sep 17 00:00:00 2001 From: saphid <4596216+saphid@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:25:41 +1000 Subject: [PATCH] Devices: fixes from review round 1 - No switching headsets (or changing the active one's user/port, or removing it) while installs run: they read the ssh settings step by step. - Switching reroutes every command to the new headset at once, even if it never answers. - ~/.ssh/config edits are serialized, use unique temp files, and back off if another program wrote the file meanwhile. - stop() ends a handshake in progress and joins the connector. - Pinned keys are written unhashed (HashKnownHosts=no); hashed ones are still found and forgotten via ssh-keygen. - Set Up Connection changing a headset's user or port updates the registry. Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/fakessh/ssh | 5 +- tests/test_devices.py | 31 +++++++++++ tests/test_link.py | 43 ++++++++++++++ ui/frame_devices.py | 127 +++++++++++++++++++++++++++++------------- ui/frame_link.py | 39 ++++++++++++- ui/server.py | 58 ++++++++++++++++--- 6 files changed, 252 insertions(+), 51 deletions(-) diff --git a/tests/fakessh/ssh b/tests/fakessh/ssh index af0ebb4..45b2b17 100755 --- a/tests/fakessh/ssh +++ b/tests/fakessh/ssh @@ -1,7 +1,8 @@ #!/usr/bin/env python3 """A stand-in for OpenSSH's ssh, for tests/test_link.py: prints what `ssh -v` prints at each step of a connection, and plays a ControlMaster. What each HostName does comes -from $FAKESSH_HOSTS (JSON: host -> "ok", "wrong" (a different host key) or "denied"); +from $FAKESSH_HOSTS (JSON: host -> "ok", "wrong" (a different host key), "denied" or +"slow" (hangs after connecting); every call is appended to $FAKESSH_LOG as a JSON line. POSIX only.""" import json import os @@ -49,6 +50,8 @@ if what is None: sys.exit(255) say(f"debug1: Connecting to {host} [127.0.0.1] port {opts.get('port', 22)}.") say("debug1: Connection established.") +if what == "slow": + time.sleep(30) say(f"debug1: Authenticating to {host}:22 as '{opts.get('user', 'tester')}'") say("debug1: Server host key: ssh-ed25519 SHA256:fakefakefakefakefakefakefakefakefakefakefak") if what == "wrong": diff --git a/tests/test_devices.py b/tests/test_devices.py index 83c187c..a115071 100644 --- a/tests/test_devices.py +++ b/tests/test_devices.py @@ -120,6 +120,13 @@ class Migration(Base): hosts = [a["host"] for a in self.reg.by_alias("frame")["addresses"]] self.assertEqual(hosts, ["192.168.1.237", "frame.tail1234.ts.net"]) # the new one first + def test_setup_changing_the_login_updates_the_headset(self): + self.reg.sync_from_config(seed=False) + (self.ssh / "config").write_text(CONFIG.replace(" User steamos\n", " User deck\n Port 2200\n", 1)) + self.assertTrue(self.reg.sync_from_config(seed=False)) + d = self.reg.by_alias("frame") + self.assertEqual((d["user"], d["port"]), ("deck", 2200)) + def test_removed_headset_stays_removed_until_setup_changes_it(self): self.reg.sync_from_config(seed=False) second = self.reg.by_alias("frame-2") @@ -159,6 +166,21 @@ class ConfigRewrite(Base): if os.name != "nt": self.assertEqual(cfg.stat().st_mode & 0o777, 0o600) + def test_concurrent_edits_all_land(self): + import threading + def edit(alias, prefix): + for n in range(15): + fd.rewrite_block(alias, hostname=f"{prefix}.{n}") + threads = [threading.Thread(target=edit, args=("frame", "10.0.0")), + threading.Thread(target=edit, args=("frame-2", "10.0.1"))] + for t in threads: + t.start() + for t in threads: + t.join() + blocks = fd.parse_blocks((self.ssh / "config").read_text()) + self.assertEqual([b["hostname"] for b in blocks], ["10.0.0.14", "10.0.1.14"]) + self.assertEqual([p.name for p in self.ssh.iterdir() if "frame-control." in p.name], []) # no temp files left + def test_zone_is_escaped_and_read_back(self): fd.rewrite_block("frame", hostname="fe80::1%en0") self.assertIn("HostName fe80::1%%en0", (self.ssh / "config").read_text()) @@ -193,6 +215,15 @@ class Pins(Base): self.assertTrue(fd.seed_pin("d3", ["frame.local"], port=2222)) self.assertIn(f"frame-control-d3 {KEY}", fd.known_hosts().read_text()) + def test_hashed_pins_are_found_and_forgotten(self): + fd.known_hosts().write_text(f"frame-control-d4 {KEY}\n") + subprocess.run(["ssh-keygen", "-H", "-f", str(fd.known_hosts())], capture_output=True, check=True) + self.assertNotIn("frame-control-d4", fd.known_hosts().read_text()) + self.assertTrue(fd.pinned("d4")) + self.assertTrue(fd.forget_pin("d4")) + self.assertFalse(fd.pinned("d4")) + self.assertFalse(fd.known_hosts().with_name("frame-control_known_hosts.old").exists()) + def test_known_hosts_option_uses_the_override(self): self.assertEqual(fd.known_hosts_opt(), str(self.ssh / "frame-control_known_hosts")) diff --git a/tests/test_link.py b/tests/test_link.py index a2575b4..c2ea013 100644 --- a/tests/test_link.py +++ b/tests/test_link.py @@ -227,6 +227,49 @@ class Connecting(unittest.TestCase): self.assertEqual(self.routes, []) self.assertTrue(all("ControlPath=none" in c for c in self.calls())) + def test_switching_to_a_headset_that_never_answers_stops_using_the_last_one(self): + self.device("localhost") + self.hosts({"localhost": "ok"}) + self.link.connect(["start"]) + other = self.reg.add_device("frame-other", port=self.port) + self.reg.add_address(other["id"], "nothing.invalid") + self.link.use(other["id"]) + self.link.connect(["switch"]) + self.assertEqual(self.link.snapshot()["phase"], "failed") + alias, opts = self.routes[-1] + self.assertEqual(alias, "frame-other") + self.assertIn("HostName=nothing.invalid", opts) + self.assertIn(f"HostKeyAlias=frame-control-{other['id']}", opts) + + def test_no_switching_while_something_is_installing(self): + d = self.device("localhost") + other = self.reg.add_device("frame-other") + for body in ({"action": "use", "id": other["id"]}, {"action": "remove", "id": d["id"]}, + {"action": "update", "id": d["id"], "port": 2222}): + with self.assertRaises(fd.DeviceError, msg=body): + fl.devices_action(self.link, body, open_setup=None, busy=lambda: 1) + # Renaming, or changing another headset, is fine. + fl.devices_action(self.link, {"action": "update", "id": d["id"], "name": "Desk"}, None, busy=lambda: 1) + fl.devices_action(self.link, {"action": "update", "id": other["id"], "port": 2222}, None, busy=lambda: 1) + self.assertEqual(self.reg.get(d["id"])["name"], "Desk") + + def test_stopping_mid_handshake_leaves_no_ssh_behind(self): + self.device("localhost") + self.hosts({"localhost": "slow"}) + t = threading.Thread(target=self.link.connect, args=(["start"],), daemon=True) + t.start() + for _ in range(100): + if self.link.pending: + break + time.sleep(0.05) + proc = self.link.pending + self.assertIsNotNone(proc) + self.link.stop() + t.join(10) + self.assertFalse(t.is_alive()) + self.assertIsNotNone(proc.poll()) + self.assertIsNone(self.link.master) + def test_devices_api_checks_everything(self): d = self.device("localhost") bad = [{"action": "address-add", "id": d["id"], "host": "-oProxyCommand=touch /tmp/x"}, diff --git a/ui/frame_devices.py b/ui/frame_devices.py index 3b5a517..b207ad9 100644 --- a/ui/frame_devices.py +++ b/ui/frame_devices.py @@ -33,6 +33,7 @@ import os import re import secrets import subprocess +import tempfile import threading import time from pathlib import Path @@ -178,34 +179,63 @@ def read_config(path=None): return "" -def _write_config(path, text): - """Swap the file in whole (same as frame_connect.write_config), keeping it private.""" - tmp = path.with_name("config.frame-control.tmp") - tmp.write_text(text, encoding="utf-8") - if not frame_host.WINDOWS: - tmp.chmod(0o600) - for attempt in range(20): # Windows: a running ssh.exe can hold the file for a moment - try: - os.replace(tmp, path) - return - except PermissionError: - time.sleep(0.25) - tmp.unlink() - raise OSError(f"{path} stayed locked by another program") +# One edit of ~/.ssh/config at a time from this app (the connector and the page can +# both want one); _edit_config also notices another program writing in between. +_config_lock = threading.Lock() + + +def _write_config(path, text, expected): + """Swap the file in whole (as frame_connect.write_config does), keeping it private. + Returns False, writing nothing, if the file no longer holds `expected`.""" + fd_, tmp = tempfile.mkstemp(prefix="config.frame-control.", dir=str(path.parent)) + tmp = Path(tmp) + try: + with os.fdopen(fd_, "w", encoding="utf-8") as fh: + fh.write(text) + if not frame_host.WINDOWS: + tmp.chmod(0o600) + for attempt in range(20): # Windows: a running ssh.exe can hold the file for a moment + if read_config(path) != expected: + return False + try: + os.replace(tmp, path) + return True + except PermissionError: + time.sleep(0.25) + raise OSError(f"{path} stayed locked by another program") + finally: + if tmp.exists(): + tmp.unlink() + + +def _edit_config(path, change): + """Apply change(lines) -> new lines or None to the file, retrying if another program + wrote it meanwhile. -> True if the file changed.""" + with _config_lock: + for _ in range(5): + text = read_config(path) + new = change(text.splitlines()) + if new is None: + return False + if _write_config(path, "\n".join(new) + "\n", text): + return True + raise OSError(f"{path} kept changing while Frame Control tried to update it") def rewrite_block(alias, path=None, hostname=None, user=None, port=None): """Change HostName, User or Port inside ALIAS's managed block, leaving the rest of the file alone. -> True if the file changed. Does nothing if there's no such block.""" - path = Path(path or ssh_config()) - text = read_config(path) - lines = text.splitlines() + return _edit_config(Path(path or ssh_config()), + lambda lines: _rewritten(lines, alias, hostname, user, port)) + + +def _rewritten(lines, alias, hostname, user, port): begin, end = begin_mark(alias), end_mark(alias) if begin not in lines or end not in lines: - return False + return None i, j = lines.index(begin), lines.index(end) if j < i: - return False + return None block = lines[i:j] want = {"hostname": ssh_host(hostname) if hostname else None, "user": user, "port": str(port) if port else None} @@ -224,23 +254,16 @@ def rewrite_block(alias, path=None, hostname=None, user=None, port=None): at = next((n + 1 for n, line in enumerate(out) if line.split(None, 1)[:1] == ["HostName"]), 2) out.insert(at, f" Port {want['port']}") new = lines[:i] + out + lines[j:] - if new == lines: - return False - _write_config(path, "\n".join(new) + "\n") - return True + return None if new == lines else new def remove_block(alias, path=None): - path = Path(path or ssh_config()) - lines = read_config(path).splitlines() - begin, end = begin_mark(alias), end_mark(alias) - if begin not in lines or end not in lines: - return False - i, j = lines.index(begin), lines.index(end) - if j < i: - return False - _write_config(path, "\n".join(lines[:i] + lines[j + 1:]) + "\n") - return True + def change(lines): + begin, end = begin_mark(alias), end_mark(alias) + if begin not in lines or end not in lines or lines.index(end) < lines.index(begin): + return None + return lines[:lines.index(begin)] + lines[lines.index(end) + 1:] + return _edit_config(Path(path or ssh_config()), change) # ---- pinned host keys ------------------------------------------------------------- @@ -252,10 +275,26 @@ def _pin_lines(path=None): return [] +def _keygen(*args): + try: + return subprocess.run(["ssh-keygen", *args], capture_output=True, stdin=subprocess.DEVNULL, text=True, + timeout=10) + except (OSError, subprocess.TimeoutExpired): + return None + + def pinned(device_id, path=None): + """Whether a key is saved for the device. The app's own entries are plain text (it + passes HashKnownHosts=no), but ask ssh-keygen too in case one was hashed.""" name = host_key_alias(device_id) - return any(line.split(None, 1)[0].split(",").count(name) for line in _pin_lines(path) - if line.strip() and not line.startswith("#")) + target = Path(path or known_hosts()) + if any(line.split(None, 1)[0].split(",").count(name) for line in _pin_lines(target) + if line.strip() and not line.startswith("#")): + return True + if not target.is_file(): + return False + r = _keygen("-F", name, "-f", str(target)) + return bool(r and r.returncode == 0 and r.stdout.strip()) def seed_pin(device_id, hosts, port=22, sources=None, path=None): @@ -299,10 +338,16 @@ def forget_pin(device_id, path=None): name = host_key_alias(device_id) lines = _pin_lines(target) kept = [line for line in lines if not (line.strip() and name in line.split(None, 1)[0].split(","))] - if kept != lines: + removed = kept != lines + if removed: target.write_text("".join(line + "\n" for line in kept), encoding="utf-8") - return True - return False + if target.is_file() and pinned(device_id, target): # a hashed entry: ssh-keygen finds it + r = _keygen("-R", name, "-f", str(target)) + removed = removed or bool(r and r.returncode == 0) + old = target.with_name(target.name + ".old") # ssh-keygen -R leaves a backup + if old.exists(): + old.unlink() + return removed # ---- address order -------------------------------------------------------------------- @@ -608,6 +653,12 @@ class Registry: if seed: seed_pin(d["id"], [host], b["port"]) changed = True + if d.get("config_login") != [user, b["port"]]: + # Set Up Connection (or an edit) changed who to log in as, or the port. + if d.get("config_login") is not None and [d["user"], d["port"]] != [user, b["port"]]: + d["user"], d["port"] = user, b["port"] if 1 <= b["port"] <= 65535 else d["port"] + d["config_login"] = [user, b["port"]] + changed = True if d["identity_files"] != b["identity_files"] and b["identity_files"]: d["identity_files"] = b["identity_files"][:8] changed = True diff --git a/ui/frame_link.py b/ui/frame_link.py index 79bc722..f23fa02 100644 --- a/ui/frame_link.py +++ b/ui/frame_link.py @@ -152,6 +152,8 @@ class Link: self.last_attempt = 0 self.config_mtime = None self.thread = None + self.pending = None # an ssh handshake still running + self.routed = None # the device id every ssh command points at # ---- publishing ---- def publish(self, **fields): @@ -218,6 +220,9 @@ class Link: self.stopped = True self.cond.notify_all() self.close_master() + if self.thread: + self.thread.join(5) # an attempt in progress notices `stopped` and ends + self.close_master() def alive(self): if self.state["phase"] != "connected": @@ -278,7 +283,7 @@ class Link: return [] return ["-o", f"HostName={frame_devices.ssh_host(host)}", "-o", f"HostKeyAlias={frame_devices.host_key_alias(device['id'])}", - "-o", f"UserKnownHostsFile={frame_devices.known_hosts_opt()}", + "-o", f"UserKnownHostsFile={frame_devices.known_hosts_opt()}", "-o", "HashKnownHosts=no", "-o", f"User={device['user']}", "-o", f"Port={device['port']}"] def public_device(self, d): @@ -333,8 +338,12 @@ class Link: mtime = None if mtime != self.config_mtime: self.config_mtime = mtime + before = self.active_device() if self.reg.sync_from_config(): self.devices_changed() + after = self.active_device() + if (before.get("user"), before.get("port")) != (after.get("user"), after.get("port")): + self.kick("switch") # Set Up Connection changed the active headset's login def refresh_network(self): net = frame_network.current_network(self.last_fp) @@ -348,6 +357,13 @@ class Link: self.last_attempt = now() self.close_master() device = self.active_device() + if device["id"] != self.routed: + # Another headset: nothing may go on reaching the last one, even if this one + # never answers. Its first address (and its own pinned identity) until one does. + first = device["addresses"][0]["host"] if device["addresses"] else None + self.alias, self.opts = device["alias"], self.host_opts(device, first) + self.apply(self.alias, self.opts) + self.routed = device["id"] with self.cond: self.state.update(phase="connecting", reason=why, device=self.public_device(device), via=None, error=None, retry_at=None, attempt=self.state["attempt"] + 1, started=now(), @@ -541,6 +557,9 @@ class Link: def close_master(self): proc, self.master = self.master, None + pending, self.pending = self.pending, None + if pending and pending.poll() is None: + pending.kill() if self.control and self.alias: try: subprocess.run([*self.mux_base, *self.opts, "-O", "exit", self.alias], capture_output=True, @@ -587,6 +606,7 @@ class Link: except OSError as e: self.fail("ssh", f"Couldn't run ssh: {e}", str(e)) return "stop" + self.pending = proc # so stop() can end it mid-handshake lines = queue.Queue() collecting = [True] @@ -603,6 +623,9 @@ class Link: mismatch = False while True: left = deadline - time.monotonic() + if self.stopped: + proc.kill() + return "stop" if left <= 0: proc.kill() self.fail(step, self.explain(f"Timed out talking to {alias}") or "The headset took too long to answer.", @@ -659,6 +682,10 @@ class Link: self.stage("login", "done", f"Logged in as {user}") step = "connected" collecting[0] = False # the master keeps printing mux debug lines: drop them + self.pending = None + if self.stopped: + proc.kill() + return "stop" for sid in ("ssh", "identity", "login"): with self.cond: pending = any(s["id"] == sid and s["state"] != "done" for s in self.state["stages"]) @@ -761,13 +788,19 @@ def devices_view(link): "kinds": frame_devices.KIND_LABEL} -def devices_action(link, body, open_setup): - """POST /api/devices {"action": ..., "id": device id, ...}. -> {"message", ...devices_view}.""" +def devices_action(link, body, open_setup, busy=lambda: 0): + """POST /api/devices {"action": ..., "id": device id, ...}. -> {"message", ...devices_view}. + busy() counts installs in progress: nothing may move them to another headset.""" reg = link.reg action = body.get("action") did = body.get("id") active = link.active_device() is_active = did == active["id"] + moves = action == "use" or (is_active and (action == "remove" or ( + action == "update" and (body.get("user") is not None or body.get("port") is not None)))) + if moves and busy(): + raise frame_devices.DeviceError( + f"Wait for what's running on {active['name']} to finish (see the activity bar), then try again") if action == "use": d = reg.get(did) link.use(did) diff --git a/ui/server.py b/ui/server.py index e578ae4..949d300 100755 --- a/ui/server.py +++ b/ui/server.py @@ -14,6 +14,7 @@ Env: FRAME_ALIAS (default frame) """ import argparse import base64 +import contextlib import http.client import json import os @@ -77,20 +78,53 @@ HOST_OPTS = [] frame_android.SSH_OPTS = SSH[1:] +_route_lock = threading.Lock() + + def route(alias, host_opts): """Point every ssh, scp and rsync at `alias` with `host_opts` (frame_link calls this when it picks a headset and an address). The lists change in place, so code holding them follows; frame_titles reads frame_android.SSH_OPTS at call time.""" global FRAME, HOST_OPTS - FRAME = frame_android.FRAME = alias - HOST_OPTS = list(host_opts) - MUX[:] = [*MUX_BASE, *HOST_OPTS] - SSH[:] = [*MUX, *SSH_TAIL] - frame_android.SSH_OPTS = SSH[1:] + with _route_lock: + FRAME = frame_android.FRAME = alias + HOST_OPTS = list(host_opts) + MUX[:] = [*MUX_BASE, *HOST_OPTS] + SSH[:] = [*MUX, *SSH_TAIL] + frame_android.SSH_OPTS = SSH[1:] LINK = None # the connector (frame_link.Link); None on the Frame itself +# Installs and other changes in progress. Switching headsets waits for them: they +# read the ssh settings step by step, so a switch could send the rest (or a failed +# install's clean-up) to the other headset. +_work_lock = threading.Lock() +_work = [0] + + +@contextlib.contextmanager +def working(): + with _work_lock: + _work[0] += 1 + try: + yield + finally: + with _work_lock: + _work[0] -= 1 + + +def busy_while(fn): + def run(*args, **kwargs): + with working(): + return fn(*args, **kwargs) + return run + + +def busy(): + with _work_lock: + return _work[0] + APPID = re.compile(r"^\d{1,10}$") FLATPAK_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]*(\.[A-Za-z0-9_-]+){2,}$") MAX_UPLOAD = 8 * 1024**3 @@ -181,7 +215,8 @@ def start_job(label, work): def run(): fields = {} try: - result = work() + with working(): + result = work() fields = {"message": result.get("message") or f"{label}: done", "result": result} except (Failure, frame_android.FrameError) as e: fields = {"error": unreachable(str(e)) or str(e)} @@ -656,6 +691,7 @@ def stage_title(path, temp_dir=None, name=None): "token": token, "plan": frame_titles.public(plan)} +@busy_while def _run_title_install(token, entry, name, exe, runtime): def update(**fields): # the page reads jobs from other threads; change them under the lock with _titles_lock: @@ -1090,6 +1126,7 @@ def webinstall_start(body): return {"job": pid} +@busy_while def _webinstall_run(plan, job): tmp = None try: @@ -1244,7 +1281,7 @@ def devices_post(body): if not LINK: raise Failure("Headsets are managed from the computer app", 400) try: - return frame_link.devices_action(LINK, body, open_setup) + return frame_link.devices_action(LINK, body, open_setup, busy) except frame_devices.DeviceError as e: raise Failure(str(e), 400) @@ -1431,7 +1468,8 @@ class Handler(BaseHTTPRequestHandler): path = urlparse(self.path).path try: if path == "/api/upload": - self.send_json(self.upload()) + with working(): + self.send_json(self.upload()) return handler = POST.get(path) if not handler: @@ -1443,7 +1481,9 @@ class Handler(BaseHTTPRequestHandler): body = json.loads(self.rfile.read(length) or b"{}") if not isinstance(body, dict): raise Failure("request body must be a JSON object", 400) - self.send_json(handler(body)) + with (contextlib.nullcontext() if path == "/api/devices" else working()): + result = handler(body) + self.send_json(result) except Failure as e: self.send_error_json(str(e), e.status, e.apk) except (ValueError, TypeError) as e: