feat(operations): add connectivity, config and lifecycle operations
This commit is contained in:
@@ -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)
|
||||
@@ -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
|
||||
@@ -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 " "
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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": ""}
|
||||
Reference in New Issue
Block a user