feat(operations): add connectivity, config and lifecycle operations

This commit is contained in:
ahp
2026-09-12 13:47:31 +03:30
parent 60ba7e8286
commit e999ba0e19
8 changed files with 1328 additions and 0 deletions
View File
+107
View File
@@ -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)
+44
View File
@@ -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
+471
View File
@@ -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 " "
+62
View File
@@ -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
+86
View File
@@ -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)
+331
View File
@@ -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,
)
+227
View File
@@ -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": ""}