From e999ba0e19ced9531960e742882fe6e0e04522a2 Mon Sep 17 00:00:00 2001 From: Amir Husayn Panahifar Date: Sat, 12 Sep 2026 13:47:31 +0330 Subject: [PATCH] feat(operations): add connectivity, config and lifecycle operations --- src/sshctl/operations/__init__.py | 0 src/sshctl/operations/checker.py | 107 +++++++ src/sshctl/operations/cleanup.py | 44 +++ src/sshctl/operations/configops.py | 471 +++++++++++++++++++++++++++++ src/sshctl/operations/events.py | 62 ++++ src/sshctl/operations/executor.py | 86 ++++++ src/sshctl/operations/jump.py | 331 ++++++++++++++++++++ src/sshctl/operations/setup.py | 227 ++++++++++++++ 8 files changed, 1328 insertions(+) create mode 100644 src/sshctl/operations/__init__.py create mode 100644 src/sshctl/operations/checker.py create mode 100644 src/sshctl/operations/cleanup.py create mode 100644 src/sshctl/operations/configops.py create mode 100644 src/sshctl/operations/events.py create mode 100644 src/sshctl/operations/executor.py create mode 100644 src/sshctl/operations/jump.py create mode 100644 src/sshctl/operations/setup.py diff --git a/src/sshctl/operations/__init__.py b/src/sshctl/operations/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/sshctl/operations/checker.py b/src/sshctl/operations/checker.py new file mode 100644 index 0000000..be1fdce --- /dev/null +++ b/src/sshctl/operations/checker.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import asyncio +import time +from dataclasses import dataclass +from pathlib import Path + +from sshctl.config.resolver import ResolvedHost, resolve_all_hosts + + +@dataclass +class HostCheck: + alias: str + hostname: str + user: str + port: int + proxy_jump: str + dns_ok: bool = False + dns_error: str = "" + tcp_ok: bool = False + tcp_error: str = "" + ssh_ok: bool = False + ssh_error: str = "" + rtt_ms: float = 0.0 + + @property + def passed(self) -> bool: + return self.dns_ok and self.tcp_ok + + @property + def summary(self) -> str: + dns = "DNS" if self.dns_ok else "DNS!" + tcp = "TCP" if self.tcp_ok else "TCP!" + ssh = "SSH" if self.ssh_ok else "SSH!" if self.ssh_error else "---" + return "/".join([dns, tcp, ssh]) + + +async def check_host( + host: ResolvedHost, + *, + timeout: int = 10, +) -> HostCheck: + check = HostCheck( + alias=host.alias, + hostname=host.hostname, + user=host.user, + port=host.port, + proxy_jump=host.proxy_jump, + ) + + start = time.monotonic() + + try: + await asyncio.wait_for( + _check_dns(host.hostname), + timeout=timeout, + ) + check.dns_ok = True + except (OSError, TimeoutError) as e: + check.dns_error = str(e) + + if check.dns_ok: + try: + await asyncio.wait_for( + _check_tcp(host.hostname, host.port), + timeout=timeout, + ) + check.tcp_ok = True + except (OSError, TimeoutError) as e: + check.tcp_error = str(e) + + check.rtt_ms = (time.monotonic() - start) * 1000 + return check + + +async def _check_dns(hostname: str) -> None: + await asyncio.get_event_loop().getaddrinfo(hostname, None) + + +async def _check_tcp(hostname: str, port: int) -> None: + _, writer = await asyncio.open_connection(hostname, port) + writer.close() + await writer.wait_closed() + + +async def check_all_hosts( + config_path: str | Path, + *, + timeout: int = 10, + max_concurrent: int = 50, + filter_alias: str | None = None, +) -> list[HostCheck]: + hosts = resolve_all_hosts(config_path) + + if filter_alias: + from fnmatch import fnmatch + + hosts = [h for h in hosts if fnmatch(h.alias, filter_alias)] + + sem = asyncio.Semaphore(max_concurrent) + + async def _checked(host: ResolvedHost) -> HostCheck: + async with sem: + return await check_host(host, timeout=timeout) + + results = await asyncio.gather(*[_checked(h) for h in hosts]) + return sorted(results, key=lambda r: r.alias) diff --git a/src/sshctl/operations/cleanup.py b/src/sshctl/operations/cleanup.py new file mode 100644 index 0000000..62eb2c3 --- /dev/null +++ b/src/sshctl/operations/cleanup.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import subprocess +from pathlib import Path + + +def cleanup_control_sockets(ssh_dir: str | Path, dry_run: bool = False) -> dict[str, int]: + """Remove stale SSH control master sockets. + + Returns dict with 'total', 'removed', 'active' counts. + """ + socket_dir = Path(ssh_dir) / "control.d" + if not socket_dir.is_dir(): + return {"total": 0, "removed": 0, "active": 0} + + total = 0 + removed = 0 + active = 0 + + for sock in sorted(socket_dir.glob("*.sock")): + if not sock.is_file(): + continue + total += 1 + if _is_socket_in_use(sock): + active += 1 + else: + if not dry_run: + sock.unlink() + removed += 1 + + return {"total": total, "removed": removed, "active": active} + + +def _is_socket_in_use(sock: Path) -> bool: + try: + result = subprocess.run( + ["lsof", "-n", str(sock)], + capture_output=True, + text=True, + timeout=5, + ) + return "ssh" in result.stdout + except (FileNotFoundError, subprocess.TimeoutExpired, OSError): + return False diff --git a/src/sshctl/operations/configops.py b/src/sshctl/operations/configops.py new file mode 100644 index 0000000..eedee95 --- /dev/null +++ b/src/sshctl/operations/configops.py @@ -0,0 +1,471 @@ +from __future__ import annotations + +import re +from pathlib import Path + +_HOST_RE = re.compile(r"^Host\s+(.+)$", re.I) +_OPTION_RE = re.compile(r"^(\s+)([a-zA-Z0-9]+)\s+(.+)$") +_INCLUDE_RE = re.compile(r"^Include\s+", re.I) +_COMMENT_RE = re.compile(r"^\s*(#|$)") + + +_VALID_OPTIONS = { + "addressfamily", + "batchmode", + "bindaddress", + "bindinterface", + "canonicaldomains", + "canonicalizefallbacklocal", + "canonicalizehostname", + "canonicalizemaxdots", + "canonicalizename", + "canonicalizepermittedcnames", + "certificatefile", + "checkhostip", + "ciphers", + "compression", + "compressionlevel", + "connectionattempts", + "connecttimeout", + "controlmaster", + "controlpath", + "controlpersist", + "dynamicforward", + "enableescapecommandline", + "escapechar", + "exitonforwardfailure", + "fingerprinthash", + "forkafterauthentication", + "forwardagent", + "forwardx11", + "forwardx11timeout", + "forwardx11trusted", + "gatewayports", + "globalknownhostsfile", + "gssapiauthentication", + "gssapidelegatecredentials", + "gssapikeyexchange", + "gssapirenewalforcesrekey", + "gssapitrustdns", + "hashknownhosts", + "hostbasedacceptedalgorithms", + "hostbasedauthentication", + "hostkeyalgorithms", + "hostkeyalias", + "hostname", + "identitiesonly", + "identityagent", + "identityfile", + "ignoreunknown", + "include", + "ipqos", + "kexalgorithms", + "knownhostscommand", + "localcommand", + "localforward", + "loglevel", + "macs", + "match", + "masters enabled", + "passwordauthentication", + "path", + "permitlocalcommand", + "permittunnel", + "pkcs11provider", + "port", + "preferredauthentications", + "protocol", + "proxycommand", + "proxyjump", + "pubkeyacceptedalgorithms", + "pubkeyauthentication", + "rekeylimit", + "remotecommand", + "remote forward", + "requesttty", + "requiredrsasize", + "revokedhostkeys", + "securitykeyprovider", + "sendenv", + "serveralivecountmax", + "serveraliveinterval", + "sessiontype", + "setenv", + "stricthostkeychecking", + "syslogfacility", + "tcpkeepalive", + "tunnel", + "tunneldevice", + "updatehostkeys", + "useroaming", + "user", + "userknownhostsfile", + "verifyhostkeydns", + "visualhostkey", + "xauthlocation", +} + +_KNOWN_TYPOS = { + "hostname": "HostName", + "host": "HostName", + "user": "User", + "port": "Port", + "identityfile": "IdentityFile", + "proxyjump": "ProxyJump", + "proxycommand": "ProxyCommand", + "strict hostkeychecking": "StrictHostKeyChecking", + "stricthostkey checking": "StrictHostKeyChecking", + "hostkeyalgorithms": "HostKeyAlgorithms", + "kexalgorithms": "KexAlgorithms", + "preferredauthentications": "PreferredAuthentications", + "ciphers": "Ciphers", + "forwardagent": "ForwardAgent", + "connecttimeout": "ConnectTimeout", +} + +_COMMON_INVALID = { + "gssapikeyexchange", # not supported by all builds +} + + +def get_host_star_block(config_path: Path) -> dict[str, str]: + """Parse just the Host * block from the main config.""" + options: dict[str, str] = {} + in_host_star = False + + for line in config_path.read_text().splitlines(keepends=False): + stripped = line.strip() + + m = _HOST_RE.match(stripped) + if m: + aliases = m.group(1).strip() + in_host_star = aliases == "*" + continue + + if in_host_star: + if _INCLUDE_RE.match(stripped): + break + if _COMMENT_RE.match(stripped): + continue + match = _OPTION_RE.match(line) + if match: + key = match.group(2) + val = match.group(3).strip().strip('"').strip("'") + options[key] = val + + return options + + +def set_host_star_option(config_path: Path, key: str, value: str) -> dict[str, str]: + """Set an option in the Host * block. Creates the option if missing.""" + lines = config_path.read_text().splitlines(keepends=False) + result: list[str] = [] + in_host_star = False + found = False + inserted = False + + for i, line in enumerate(lines): + stripped = line.strip() + + m = _HOST_RE.match(stripped) + if m: + aliases = m.group(1).strip() + in_host_star = aliases == "*" + result.append(line) + continue + + if in_host_star: + if _INCLUDE_RE.match(stripped): + if not found and not inserted: + indent = _detect_indent(lines, i) or " " + result.append(f"{indent}{key} {value}") + inserted = True + in_host_star = False + result.append(line) + continue + + if _COMMENT_RE.match(stripped): + result.append(line) + continue + + match = _OPTION_RE.match(line) + if match and match.group(2).lower() == key.lower(): + indent = match.group(1) + result.append(f"{indent}{key} {value}") + found = True + continue + + result.append(line) + continue + + result.append(line) + + if not found and not inserted: + result = ["Host *", f" {key} {value}", ""] + result + config_path.write_text("\n".join(result) + "\n") + return {key: value} + + config_path.write_text("\n".join(result) + "\n") + return {key: value} + + +def unset_host_star_option(config_path: Path, key: str) -> bool: + """Remove an option from the Host * block.""" + lines = config_path.read_text().splitlines(keepends=False) + result: list[str] = [] + in_host_star = False + removed = False + + for line in lines: + stripped = line.strip() + + m = _HOST_RE.match(stripped) + if m: + aliases = m.group(1).strip() + in_host_star = aliases == "*" + result.append(line) + continue + + if in_host_star: + if _INCLUDE_RE.match(stripped): + in_host_star = False + result.append(line) + continue + + if _COMMENT_RE.match(stripped): + result.append(line) + continue + + match = _OPTION_RE.match(line) + if match and match.group(2).lower() == key.lower(): + removed = True + continue + result.append(line) + continue + + result.append(line) + + if removed: + config_path.write_text("\n".join(result) + "\n") + return removed + + +def validate_options(config_path: Path) -> list[dict[str, str]]: + """Check all options in the config for validity.""" + issues: list[dict[str, str]] = [] + seen_invalid: set[str] = set() + + for line in config_path.read_text().splitlines(keepends=False): + stripped = line.strip() + if _COMMENT_RE.match(stripped) or _HOST_RE.match(stripped) or _INCLUDE_RE.match(stripped): + continue + + match = _OPTION_RE.match(line) + if match: + key = match.group(2) + key_lower = key.lower() + + if key_lower in _COMMON_INVALID: + issues.append( + { + "type": "invalid", + "key": key, + "line": stripped, + "message": f"'{key}' is not supported by all OpenSSH builds", + } + ) + seen_invalid.add(key_lower) + continue + + fixed = _KNOWN_TYPOS.get(key_lower) or _KNOWN_TYPOS.get(key) + if fixed: + if key != fixed: + issues.append( + { + "type": "typo", + "key": key, + "fixed": fixed, + "line": stripped, + "message": f"'{key}' → '{fixed}'", + } + ) + elif key_lower not in _VALID_OPTIONS: + issues.append( + { + "type": "unknown", + "key": key, + "line": stripped, + "message": f"'{key}' is not a recognized OpenSSH option", + } + ) + + return issues + + +def heal_config(config_path: Path, auto_fix: bool = False) -> list[dict[str, str]]: + """Detect and fix issues in the SSH config. + + Returns list of actions taken or proposed. + """ + issues = validate_options(config_path) + actions: list[dict[str, str]] = [] + + if auto_fix: + lines = config_path.read_text().splitlines(keepends=False) + result: list[str] = [] + + for line in lines: + match = _OPTION_RE.match(line) + if match: + key = match.group(2) + key_lower = key.lower() + + if key_lower in _COMMON_INVALID: + actions.append( + { + "action": "removed", + "key": key, + "message": f"Removed invalid option '{key}'", + } + ) + continue + + if key_lower in _KNOWN_TYPOS: + fixed = _KNOWN_TYPOS[key_lower] + if key == fixed: + result.append(line) + continue + indent = match.group(1) + val = match.group(3) + result.append(f"{indent}{fixed} {val}") + actions.append( + { + "action": "fixed", + "key": key, + "fixed": fixed, + "message": f"Fixed '{key}' → '{fixed}'", + } + ) + continue + + result.append(line) + + config_path.write_text("\n".join(result) + "\n") + + remaining = validate_options(config_path) + for r in remaining: + actions.append( + { + "action": "warning", + "key": r["key"], + "message": r["message"], + } + ) + else: + for issue in issues: + actions.append( + { + "action": "detected", + "key": issue["key"], + "message": issue["message"], + } + ) + + return actions + + +def find_host_file(alias: str, ssh_dir: Path) -> Path | None: + """Find the conf.d file containing a specific host alias.""" + conf_d = ssh_dir / "conf.d" + if not conf_d.is_dir(): + return None + + for f in sorted(conf_d.rglob("*.conf")): + text = f.read_text() + for line in text.splitlines(): + m = _HOST_RE.match(line.strip()) + if m: + names = m.group(1).strip().split() + if alias in names: + return f + return None + + +def get_host_aliases(ssh_dir: Path) -> list[str]: + """Get all host aliases from conf.d files.""" + aliases: list[str] = [] + conf_d = ssh_dir / "conf.d" + if not conf_d.is_dir(): + return aliases + for f in sorted(conf_d.rglob("*.conf")): + text = f.read_text() + for line in text.splitlines(): + m = _HOST_RE.match(line.strip()) + if m: + name = m.group(1).strip() + if name != "*": + aliases.append(name) + return aliases + + +def remove_host_from_file(alias: str, file_path: Path) -> bool: + """Remove a host block from a conf.d file.""" + lines = file_path.read_text().splitlines(keepends=False) + result: list[str] = [] + in_target = False + removed = False + blank_after = False + + for line in lines: + stripped = line.strip() + m = _HOST_RE.match(stripped) + if m: + names = m.group(1).strip().split() + if alias in names: + in_target = True + removed = True + blank_after = False + continue + in_target = False + + if in_target: + if not stripped or stripped.startswith("#"): + if not blank_after: + blank_after = True + result.append(line) + continue + if not stripped: + continue + result.append(line) + continue + blank_after = False + continue + + result.append(line) + + if removed: + cleaned = _clean_blank_lines(result) + file_path.write_text("\n".join(cleaned) + "\n") + return removed + + +def _clean_blank_lines(lines: list[str]) -> list[str]: + """Remove duplicate blank lines.""" + cleaned: list[str] = [] + prev_blank = False + for line in lines: + is_blank = not line.strip() + if is_blank and prev_blank: + continue + cleaned.append(line) + prev_blank = is_blank + return cleaned + + +def _detect_indent(lines: list[str], start: int) -> str: + for i in range(max(0, start - 20), start): + if i < len(lines): + match = _OPTION_RE.match(lines[i]) + if match: + return match.group(1) + return " " diff --git a/src/sshctl/operations/events.py b/src/sshctl/operations/events.py new file mode 100644 index 0000000..1c066fa --- /dev/null +++ b/src/sshctl/operations/events.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +_EVENTS: list[dict[str, Any]] = [] + + +@dataclass +class Event: + command: str + status: str + args: dict[str, Any] = field(default_factory=dict) + result: str = "" + duration_ms: float = 0.0 + + def to_dict(self) -> dict[str, Any]: + return { + "ts": time.time(), + "cmd": self.command, + "status": self.status, + "args": self.args, + "result": self.result, + "duration_ms": round(self.duration_ms, 1), + } + + +def log_event(event: Event) -> None: + _EVENTS.append(event.to_dict()) + + +def _session_file(ssh_dir: Path) -> Path: + log_dir = ssh_dir / "report" + log_dir.mkdir(parents=True, exist_ok=True) + return log_dir / "events.ndjson" + + +def flush(ssh_dir: Path) -> None: + if not _EVENTS: + return + path = _session_file(ssh_dir) + with path.open("a") as f: + for ev in _EVENTS: + f.write(json.dumps(ev) + "\n") + _EVENTS.clear() + + +def recent(ssh_dir: Path, limit: int = 50) -> list[dict[str, Any]]: + path = _session_file(ssh_dir) + if not path.exists(): + return [] + lines = path.read_text().strip().splitlines() + events = [] + for line in lines[-limit:]: + try: + events.append(json.loads(line)) + except json.JSONDecodeError: + continue + return events diff --git a/src/sshctl/operations/executor.py b/src/sshctl/operations/executor.py new file mode 100644 index 0000000..a570aeb --- /dev/null +++ b/src/sshctl/operations/executor.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import asyncio +import shlex +from dataclasses import dataclass +from pathlib import Path + +from sshctl.config.parser import resolve_includes + + +@dataclass +class ExecResult: + alias: str + returncode: int | None + stdout: str = "" + stderr: str = "" + error: str = "" + + +async def run_on_host( + alias: str, + command: str, + config_path: str | Path, + *, + ssh_bin: str = "ssh", + timeout: int = 120, +) -> ExecResult: + ssh_cmd = [ + ssh_bin, + "-F", + str(config_path), + "-o", + "ConnectTimeout=10", + "-o", + "StrictHostKeyChecking=accept-new", + alias, + "--", + ] + shlex.split(command) + + try: + proc = await asyncio.create_subprocess_exec( + *ssh_cmd, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + stdout, stderr = await asyncio.wait_for( + proc.communicate(), + timeout=timeout, + ) + return ExecResult( + alias=alias, + returncode=proc.returncode, + stdout=stdout.decode().strip(), + stderr=stderr.decode().strip(), + ) + except TimeoutError as e: + return ExecResult(alias=alias, returncode=None, error=f"Timeout: {e}") + except FileNotFoundError as e: + return ExecResult(alias=alias, returncode=None, error=f"SSH not found: {e}") + except OSError as e: + return ExecResult(alias=alias, returncode=None, error=str(e)) + + +async def run_on_all_hosts( + command: str, + config_path: str | Path, + *, + max_concurrent: int = 20, + filter_alias: str | None = None, +) -> list[ExecResult]: + config = resolve_includes(config_path) + aliases = config.aliases + + if filter_alias: + from fnmatch import fnmatch + + aliases = [a for a in aliases if fnmatch(a, filter_alias)] + + sem = asyncio.Semaphore(max_concurrent) + + async def _run(alias: str) -> ExecResult: + async with sem: + return await run_on_host(alias, command, config_path) + + results = await asyncio.gather(*[_run(a) for a in aliases]) + return sorted(results, key=lambda r: r.alias) diff --git a/src/sshctl/operations/jump.py b/src/sshctl/operations/jump.py new file mode 100644 index 0000000..49868fe --- /dev/null +++ b/src/sshctl/operations/jump.py @@ -0,0 +1,331 @@ +from __future__ import annotations + +import json +import re +import socket +import subprocess +from pathlib import Path +from typing import Any, TypedDict + +_PROXYJUMP_RE = re.compile(r"^\s+ProxyJump\s+", re.I) +_SAVE_FILE = ".proxyjump_save" +_INTERNAL_SSIDS_FILE = ".internal_ssids" +DEFAULT_INTERNAL_PATTERNS = ["MobinNet*", "Mobin*"] +_PROBE_TIMEOUT = 3 + + +class ProbeResult(TypedDict): + gateway: str + ip: str | None + port: int + reachable: bool + error: str + + +class JumpRecord(TypedDict): + file: str + line: int + text: str + + +class JumpNoopResult(TypedDict): + status: str + count: int + message: str + + +class JumpOkResult(TypedDict): + status: str + count: int + files: list[str] + + +class JumpStatusResult(TypedDict): + hosts: list[dict[str, str]] + count: int + saved: bool + + +JumpResult = JumpNoopResult | JumpOkResult + + +def _find_conf_files(ssh_dir: Path) -> list[Path]: + conf_d = ssh_dir / "conf.d" + if not conf_d.is_dir(): + return [] + return sorted(conf_d.rglob("*.conf")) + + +def _ssh_config_path(ssh_dir: Path) -> Path: + return ssh_dir / "config" + + +def _resolve_host(hostname: str, ssh_dir: Path) -> tuple[str, int] | None: + """Resolve a host alias to (IP, port) via SSH config, then DNS.""" + port = 22 + try: + from sshctl.config.parser import resolve_includes + + resolved = resolve_includes(_ssh_config_path(ssh_dir)) + for block in resolved.blocks: + if hostname in (block.alias or "").split(): + hostname = block.hostname or hostname + port = int(block.port) if block.port else 22 + break + except Exception: + pass + try: + ip = socket.gethostbyname(hostname) + return (ip, port) + except socket.gaierror: + return None + + +def _jump_gateways(ssh_dir: Path) -> set[str]: + """Collect unique jump gateway aliases/hostnames from ProxyJump directives.""" + gateways: set[str] = set() + for f in _find_conf_files(ssh_dir): + for line in f.read_text().splitlines(): + m = _PROXYJUMP_RE.match(line) + if m: + target = line.strip().split(None, 1)[1].strip() + for gw in target.split(): + gw = gw.split(":")[0] + gateways.add(gw) + return gateways + + +def _subnet_reachable(ssh_dir: Path) -> list[ProbeResult]: + """Probe each unique jump gateway. Return list of reachable gateways.""" + results: list[ProbeResult] = [] + for gw in sorted(_jump_gateways(ssh_dir)): + resolved = _resolve_host(gw, ssh_dir) + if resolved is None: + results.append( + ProbeResult( + gateway=gw, + ip=None, + port=22, + reachable=False, + error="DNS resolution failed", + ) + ) + continue + ip, port = resolved + try: + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + sock.settimeout(_PROBE_TIMEOUT) + code = sock.connect_ex((ip, port)) + sock.close() + reachable = code == 0 + error = "" if reachable else f"connect refused (port {port}, code={code})" + results.append( + ProbeResult( + gateway=gw, + ip=ip, + port=port, + reachable=reachable, + error=error, + ) + ) + except OSError as e: + results.append( + ProbeResult( + gateway=gw, + ip=ip, + port=port, + reachable=False, + error=str(e), + ) + ) + return results + + +def current_ssid() -> str | None: + try: + result = subprocess.run( + ["nmcli", "-t", "-f", "ACTIVE,SSID", "dev", "wifi"], + capture_output=True, + text=True, + timeout=5, + ) + for line in result.stdout.strip().splitlines(): + if line.startswith("yes:"): + return line[4:] + except (subprocess.SubprocessError, FileNotFoundError): + pass + return None + + +def known_networks(ssh_dir: Path) -> list[str]: + path = ssh_dir / _INTERNAL_SSIDS_FILE + if not path.exists(): + return list(DEFAULT_INTERNAL_PATTERNS) + return [ + line.strip() + for line in path.read_text().splitlines() + if line.strip() and not line.strip().startswith("#") + ] + + +def add_network(ssh_dir: Path, ssid: str) -> None: + path = ssh_dir / _INTERNAL_SSIDS_FILE + if not path.exists(): + path.write_text("\n".join(DEFAULT_INTERNAL_PATTERNS) + "\n") + existing: set[str] = set() + if path.exists(): + existing = { + line.strip() + for line in path.read_text().splitlines() + if line.strip() and not line.strip().startswith("#") + } + if ssid not in existing: + with path.open("a") as f: + f.write(ssid + "\n") + + +def is_internal_by_ssid(ssh_dir: Path) -> bool | None: + ssid = current_ssid() + if ssid is None: + return None + networks = known_networks(ssh_dir) + for pattern in networks: + if re.match(pattern.replace("*", ".*").replace("?", ".") + "$", ssid, re.I): + return True + return False + + +def is_internal_by_subnet(ssh_dir: Path) -> tuple[bool, list[ProbeResult]]: + """Check if any jump gateway is directly reachable (internal subnet).""" + probes = _subnet_reachable(ssh_dir) + reachable = any(p["reachable"] for p in probes) + return reachable, probes + + +def list_hosts(ssh_dir: Path) -> list[dict[str, str]]: + hosts: list[dict[str, str]] = [] + for f in _find_conf_files(ssh_dir): + text = f.read_text() + current_alias = "" + current_proxy = "" + for line in text.splitlines(): + stripped = line.strip() + m = re.match(r"^Host\s+(.+)$", stripped, re.I) + if m: + if current_alias and current_proxy: + hosts.append( + { + "alias": current_alias, + "jump": current_proxy, + "file": str(f.relative_to(ssh_dir)), + } + ) + current_alias = m.group(1).strip().split()[0] + current_proxy = "" + continue + pm = _PROXYJUMP_RE.match(line) + if pm and current_alias: + current_proxy = line.strip().split(None, 1)[1] + if current_alias and current_proxy: + hosts.append( + { + "alias": current_alias, + "jump": current_proxy, + "file": str(f.relative_to(ssh_dir)), + } + ) + return hosts + + +def save_state(ssh_dir: Path) -> list[JumpRecord]: + records: list[JumpRecord] = [] + for f in _find_conf_files(ssh_dir): + lines = f.read_text().splitlines(keepends=False) + for i, line in enumerate(lines): + if _PROXYJUMP_RE.match(line): + records.append( + JumpRecord( + file=str(f.relative_to(ssh_dir)), + line=i, + text=line, + ) + ) + (ssh_dir / _SAVE_FILE).write_text(json.dumps(records, indent=2)) + return records + + +def load_state(ssh_dir: Path) -> list[JumpRecord] | None: + path = ssh_dir / _SAVE_FILE + if not path.exists(): + return None + try: + data: list[JumpRecord] = json.loads(path.read_text()) + return data + except (json.JSONDecodeError, OSError): + return None + + +def off(ssh_dir: Path) -> dict[str, Any]: + records = save_state(ssh_dir) + if not records: + return {"status": "noop", "count": 0, "message": "No ProxyJump lines found"} + + modified_files: set[str] = set() + for rec in records: + modified_files.add(str(rec.get("file", ""))) + + for fname in modified_files: + fpath = ssh_dir / fname + lines = fpath.read_text().splitlines(keepends=False) + new_lines = [ln for ln in lines if not _PROXYJUMP_RE.match(ln)] + fpath.write_text("\n".join(new_lines) + "\n") + + return {"status": "ok", "count": len(records), "files": sorted(modified_files)} + + +def on(ssh_dir: Path) -> dict[str, Any]: + records = load_state(ssh_dir) + if records is None: + return { + "status": "noop", + "count": 0, + "message": "No saved state found. Run 'jump off' first.", + } + + by_file: dict[str, list[JumpRecord]] = {} + for rec in records: + fname = str(rec.get("file", "")) + by_file.setdefault(fname, []).append(rec) + + modified: set[str] = set() + for fname, file_records in by_file.items(): + fpath = ssh_dir / fname + if not fpath.exists(): + continue + file_records.sort(key=lambda r: int(r.get("line", 0)), reverse=True) + lines = fpath.read_text().splitlines(keepends=False) + if any(_PROXYJUMP_RE.match(ln) for ln in lines): + continue + for rec in file_records: + lineno = int(rec.get("line", 0)) + text = str(rec.get("text", "")) + if lineno < len(lines): + lines.insert(lineno, text) + else: + lines.append(text) + fpath.write_text("\n".join(lines) + "\n") + modified.add(fname) + + (ssh_dir / _SAVE_FILE).unlink(missing_ok=True) + + return {"status": "ok", "count": len(records), "files": sorted(modified)} + + +def status(ssh_dir: Path) -> JumpStatusResult: + hosts = list_hosts(ssh_dir) + saved = load_state(ssh_dir) + return JumpStatusResult( + hosts=hosts, + count=len(hosts), + saved=saved is not None, + ) diff --git a/src/sshctl/operations/setup.py b/src/sshctl/operations/setup.py new file mode 100644 index 0000000..40c4e80 --- /dev/null +++ b/src/sshctl/operations/setup.py @@ -0,0 +1,227 @@ +from __future__ import annotations + +import subprocess +from pathlib import Path + +SERVICE_DIR = Path("~/.config/systemd/user").expanduser() + + +def _run_systemctl(args: list[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run( + ["systemctl", "--user", *args], + capture_output=True, + text=True, + timeout=30, + ) + + +def agent_status() -> dict[str, str]: + """Check ssh-agent systemd service status.""" + result = _run_systemctl(["is-active", "ssh-agent.service"]) + return {"service": "ssh-agent", "status": result.stdout.strip() or "inactive"} + + +def agent_start() -> dict[str, str]: + """Enable and start ssh-agent systemd service.""" + _write_agent_service() + _run_systemctl(["daemon-reload"]) + _run_systemctl(["enable", "ssh-agent.service"]) + _run_systemctl(["restart", "ssh-agent.service"]) + return agent_status() + + +def agent_stop() -> dict[str, str]: + _run_systemctl(["stop", "ssh-agent.service"]) + _run_systemctl(["disable", "ssh-agent.service"]) + return {"service": "ssh-agent", "status": "stopped"} + + +def _write_agent_service() -> Path: + SERVICE_DIR.mkdir(parents=True, exist_ok=True) + path = SERVICE_DIR / "ssh-agent.service" + content = """\ +[Unit] +Description=OpenSSH agent (user session) +Documentation=man:ssh-agent(1) +After=network.target + +[Service] +Type=simple +ExecStart=/usr/bin/ssh-agent -D -a "${XDG_RUNTIME_DIR}/ssh-agent.socket" +ExecReload=/usr/bin/ssh-agent -k +ExecStop=/usr/bin/ssh-agent -k +Restart=on-failure +RestartSec=5 + +[Install] +WantedBy=default.target +""" + path.write_text(content) + return path + + +def timer_status() -> list[dict[str, str]]: + """Status of all sshctl-managed timers.""" + timers = [ + ("ssh-cleanup", "Stale SSH control socket cleanup"), + ("sshctl-check", "Weekly SSH connectivity check"), + ] + results = [] + for name, desc in timers: + result = _run_systemctl(["is-active", f"{name}.timer"]) + results.append( + { + "timer": name, + "description": desc, + "status": result.stdout.strip() or "inactive", + } + ) + return results + + +def timer_enable(name: str) -> dict[str, str]: + valid = {"cleanup": "ssh-cleanup", "check": "sshctl-check"} + service_name = valid.get(name) + if not service_name: + return {"error": f"Unknown timer '{name}'. Valid: {', '.join(valid)}"} + + _write_timer_units() + _run_systemctl(["daemon-reload"]) + _run_systemctl(["enable", f"{service_name}.timer"]) + _run_systemctl(["start", f"{service_name}.timer"]) + + result = _run_systemctl(["is-active", f"{service_name}.timer"]) + return {"timer": name, "status": result.stdout.strip() or "inactive"} + + +def timer_disable(name: str) -> dict[str, str]: + valid = {"cleanup": "ssh-cleanup", "check": "sshctl-check"} + service_name = valid.get(name) + if not service_name: + return {"error": f"Unknown timer '{name}'"} + + _run_systemctl(["stop", f"{service_name}.timer"]) + _run_systemctl(["disable", f"{service_name}.timer"]) + return {"timer": name, "status": "stopped"} + + +def _write_timer_units() -> None: + SERVICE_DIR.mkdir(parents=True, exist_ok=True) + + cleanup_service = SERVICE_DIR / "ssh-cleanup.service" + cleanup_service.write_text("""\ +[Unit] +Description=Stale SSH control socket cleanup + +[Service] +Type=oneshot +ExecStart={sshctl} cleanup +Nice=19 +IOSchedulingClass=idle +""") + + cleanup_timer = SERVICE_DIR / "ssh-cleanup.timer" + cleanup_timer.write_text("""\ +[Unit] +Description=Daily stale SSH control socket cleanup + +[Timer] +OnCalendar=daily +Persistent=true + +[Install] +WantedBy=timers.target +""") + + check_service = SERVICE_DIR / "sshctl-check.service" + check_service.write_text("""\ +[Unit] +Description=Weekly SSH connectivity check via sshctl + +[Service] +Type=oneshot +ExecStart={sshctl} check --timeout 15 +Nice=19 +IOSchedulingClass=idle +""") + + check_timer = SERVICE_DIR / "sshctl-check.timer" + check_timer.write_text("""\ +[Unit] +Description=Weekly SSH connectivity check + +[Timer] +OnCalendar=Sun 07:00 +Persistent=true + +[Install] +WantedBy=timers.target +""") + + +def init_git(ssh_dir: str | Path) -> dict[str, str]: + """Initialize git repo in ssh_dir for config tracking.""" + ssh_dir = Path(ssh_dir) + git_dir = ssh_dir / ".git" + + if git_dir.exists(): + return {"status": "exists", "path": str(ssh_dir)} + + result = subprocess.run( + ["git", "init"], + cwd=ssh_dir, + capture_output=True, + text=True, + timeout=15, + ) + if result.returncode != 0: + return {"status": "error", "message": result.stderr.strip()} + + _write_gitignore(ssh_dir) + + return {"status": "created", "path": str(ssh_dir)} + + +def _write_gitignore(ssh_dir: Path) -> None: + gitignore = ssh_dir / ".gitignore" + if gitignore.exists(): + return + gitignore.write_text("""\ +keys.d/ +keys.a/ + +keys.pub/ + +control.d/ + +known_hosts +known_hosts.old + +report/ + +inventory.d/ + +agent/ + +ssh_backup_*/ + +scripts/*.log + +.DS_Store +Thumbs.db +""") + + +def shell_hook(shell: str) -> dict[str, str]: + """Print shell hook config for SSH_AUTH_SOCK.""" + if shell == "zsh": + return { + "shell": shell, + "line": 'export SSH_AUTH_SOCK="${XDG_RUNTIME_DIR}/ssh-agent.socket"', + } + if shell == "bash": + return { + "shell": shell, + "line": 'export SSH_AUTH_SOCK="${XDG_RUNTIME_DIR}/ssh-agent.socket"', + } + return {"shell": shell, "line": ""}