Files
sshctl/src/sshctl/cli/app.py
T

1386 lines
44 KiB
Python

from __future__ import annotations
import asyncio
import re
import sys
import time
from pathlib import Path
from typing import TYPE_CHECKING, Annotated, cast
import click
import typer
from rich.console import Console
from rich.table import Table
from rich.tree import Tree
from sshctl.core import __version__
from sshctl.core.config import SshctlConfig, load_config
from sshctl.core.logutil import setup_logging
if TYPE_CHECKING:
from sshctl.core.models import HostBlock
app = typer.Typer(
name="sshctl",
help="SSH configuration management toolkit",
no_args_is_help=True,
)
console = Console()
_state = SshctlConfig()
def _version_cb(value: bool) -> None:
if value:
console.print(f"sshctl v{__version__}")
raise typer.Exit()
@app.callback()
def main(
version: Annotated[
bool | None,
typer.Option("--version", "-V", callback=_version_cb, help="Show version"),
] = None,
verbose: Annotated[bool, typer.Option("--verbose", "-v", help="Verbose output")] = False,
config_path: Annotated[
Path | None,
typer.Option("--config", "-c", help="Path to ssh_config file (default: ~/.ssh/config)"),
] = None,
) -> None:
global _state
if config_path:
yaml_path = config_path.parent / "sshctl.yaml"
_state = load_config(yaml_path) if config_path.parent else SshctlConfig()
_state.config_file = str(config_path)
else:
_state = load_config()
if verbose:
_state.verbose = True
setup_logging(verbose=True)
def _ssh_config_path() -> Path:
return _state.ssh_config_path()
def _ssh_dir() -> Path:
return _state.ssh_dir
@app.command()
def init(
ssh_dir: Annotated[
Path | None,
typer.Argument(help="SSH directory to initialize (default: ~/.ssh)"),
] = None,
) -> None:
"""Initialize an SSH config directory structure."""
target = ssh_dir.expanduser().resolve() if ssh_dir else _ssh_dir()
dirs = [
target / "conf.d",
target / "conf.d" / "internal",
target / "conf.d" / "external",
target / "control.d",
target / "keys.d",
target / "keys.pub",
target / "report",
target / "scripts",
]
for d in dirs:
d.mkdir(parents=True, exist_ok=True)
d.chmod(0o700)
(target / "control.d" / ".gitkeep").write_text("")
(target / "keys.pub" / ".gitkeep").write_text("")
config_content = (
"Host *\n"
" ## Core\n"
" User root\n"
" Port 22\n"
" ConnectTimeout 10\n"
" LogLevel ERROR\n"
" IgnoreUnknown yes\n"
"\n"
" ## Authentication\n"
' IdentityFile "~/.ssh/keys.d/id_ed25519"\n'
" IdentitiesOnly yes\n"
" PreferredAuthentications publickey\n"
" PubkeyAuthentication yes\n"
" PasswordAuthentication no\n"
" GSSAPIAuthentication no\n"
"\n"
" ## Connection multiplexing\n"
" ControlMaster auto\n"
" ControlPath ~/.ssh/control.d/%C.sock\n"
" ControlPersist 30m\n"
"\n"
" ## Keep alive\n"
" ServerAliveInterval 30\n"
" ServerAliveCountMax 3\n"
" TCPKeepAlive yes\n"
"\n"
" ## Security\n"
" StrictHostKeyChecking accept-new\n"
" UpdateHostKeys yes\n"
" HashKnownHosts yes\n"
" FingerprintHash sha256\n"
"\n"
" ## Performance\n"
" Compression no\n"
" CanonicalizeHostname no\n"
" AddressFamily inet\n"
" RekeyLimit 1G 1h\n"
"\n"
"Include conf.d/*.conf\n"
)
config_file = target / "config"
if not config_file.exists():
config_file.write_text(config_content)
config_file.chmod(0o600)
table = Table(title=f"Initialized {target}")
table.add_column("Directory/File", style="cyan")
table.add_column("Status", style="green")
for d in dirs:
rel = d.relative_to(target) if d != target else Path()
table.add_row(str(rel), "created")
table.add_row("config", "created" if config_file.exists() else "exists")
console.print(table)
config_app = typer.Typer(help="Manage global SSH config", no_args_is_help=True)
app.add_typer(config_app, name="config")
@config_app.command("show")
def config_show() -> None:
"""Show global Host * options."""
from sshctl.operations.configops import get_host_star_block
config_path = _ssh_config_path()
if not config_path.exists():
console.print("[red]Config file not found[/red]")
raise typer.Exit(1)
options = get_host_star_block(config_path)
if not options:
console.print("[yellow]No Host * block found[/yellow]")
return
table = Table(title=f"Global Options ({config_path})")
table.add_column("Option", style="cyan")
table.add_column("Value")
for key, val in sorted(options.items()):
table.add_row(key, val)
console.print(table)
@config_app.command("set")
def config_set(
key: Annotated[str, typer.Argument(help="Option name (e.g. Port)")],
value: Annotated[str, typer.Argument(help="Option value")],
) -> None:
"""Set a global Host * option."""
from sshctl.operations.configops import set_host_star_option
config_path = _ssh_config_path()
if not config_path.exists():
console.print("[red]Config file not found[/red]")
raise typer.Exit(1)
set_host_star_option(config_path, key, value)
console.print(f"[green]Set[/green] {key} = {value}")
@config_app.command("unset")
def config_unset(
key: Annotated[str, typer.Argument(help="Option name to remove")],
) -> None:
"""Remove a global Host * option."""
from sshctl.operations.configops import unset_host_star_option
config_path = _ssh_config_path()
if not config_path.exists():
console.print("[red]Config file not found[/red]")
raise typer.Exit(1)
if unset_host_star_option(config_path, key):
console.print(f"[green]Removed[/green] {key}")
else:
console.print(f"[yellow]Option '{key}' not found[/yellow]")
host_app = typer.Typer(help="Manage host aliases", no_args_is_help=True)
app.add_typer(host_app, name="host")
@host_app.command("list")
def host_list(
universe: Annotated[
str | None,
typer.Option("--universe", "-u", help="Filter by universe"),
] = None,
pattern: Annotated[
str | None,
typer.Option("--pattern", "-p", help="Glob pattern to filter aliases"),
] = None,
) -> None:
"""List all host aliases."""
from sshctl.config.parser import resolve_includes
try:
config = resolve_includes(_ssh_config_path())
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
aliases = config.aliases
if universe:
universe_conf_dir = _ssh_dir() / "conf.d" / universe
if universe_conf_dir.is_dir():
universe_aliases: list[str] = []
for conf_file in universe_conf_dir.rglob("*.conf"):
from sshctl.config.parser import parse_file
try:
parsed = parse_file(conf_file)
universe_aliases.extend(b.alias for b in parsed.blocks if b.alias != "*")
except Exception:
continue
aliases = [a for a in aliases if a in universe_aliases]
if pattern:
from fnmatch import fnmatch
aliases = [a for a in aliases if fnmatch(a, pattern)]
if not aliases:
console.print("[yellow]No hosts found[/yellow]")
raise typer.Exit()
table = Table(title=f"Hosts ({len(aliases)} total)")
table.add_column("Alias", style="cyan")
for alias in sorted(aliases):
block = config.find(alias)
if block:
hostname = block.hostname or ""
proxy = block.proxy_jump or ""
label = alias
if hostname:
label += f" -> {hostname}"
if proxy:
label += f" (via {proxy})"
table.add_row(label)
else:
table.add_row(alias)
console.print(table)
@host_app.command("show")
def host_show(
alias: Annotated[str, typer.Argument(help="Host alias to show")],
) -> None:
"""Show details for a specific host."""
from sshctl.config.parser import resolve_includes
try:
config = resolve_includes(_ssh_config_path())
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
block = config.find(alias)
if not block:
console.print(f"[red]Host '{alias}' not found[/red]")
raise typer.Exit(1)
table = Table(title=f"Host {alias}")
table.add_column("Option", style="cyan")
table.add_column("Value")
for key, val in sorted(block.options.items()):
table.add_row(key, val)
console.print(table)
@host_app.command("add")
def host_add(
alias: Annotated[str, typer.Argument(help="Host alias")],
hostname: Annotated[str, typer.Option("--hostname", "-n", help="Hostname or IP")],
user: Annotated[str | None, typer.Option("--user", "-u", help="SSH user")] = None,
port: Annotated[int | None, typer.Option("--port", "-p", help="SSH port")] = None,
proxy_jump: Annotated[str | None, typer.Option("--proxy-jump", "-j", help="Jump host")] = None,
identity_file: Annotated[
str | None, typer.Option("--identity", "-i", help="Identity file")
] = None,
) -> None:
"""Add a new host to the SSH config."""
from sshctl.config.generator import generate_host_block
from sshctl.core.models import HostBlock
host = HostBlock(alias=alias)
host.set("HostName", hostname)
if user:
host.set("User", user)
if port:
host.set("Port", str(port))
if proxy_jump:
host.set("ProxyJump", proxy_jump)
if identity_file:
host.set("IdentityFile", identity_file)
output = generate_host_block(host)
console.print(output)
conf_d = _ssh_dir() / "conf.d"
if not conf_d.exists():
console.print("[yellow]Run 'sshctl init' first[/yellow]")
raise typer.Exit(1)
out_file = conf_d / "zz-auto.conf"
with out_file.open("a") as f:
f.write(f"\n{output}\n")
console.print(f"[green]Appended to {out_file}[/green]")
@host_app.command("edit")
def host_edit(
alias: Annotated[str, typer.Argument(help="Host alias to edit")],
option: Annotated[str, typer.Option("--option", "-o", help="Option name (e.g. HostName)")],
value: Annotated[str, typer.Option("--value", "-v", help="New value")],
unset: Annotated[
str | None, typer.Option("--unset", "-u", help="Option name to remove")
] = None,
) -> None:
"""Edit or remove an option from a host."""
from sshctl.config.parser import parse_file
from sshctl.operations.configops import find_host_file
file_path = find_host_file(alias, _ssh_dir())
if not file_path:
console.print(f"[red]Host '{alias}' not found in any config file[/red]")
raise typer.Exit(1)
config = parse_file(file_path)
block = config.find(alias)
if not block:
console.print(f"[red]Host '{alias}' not found[/red]")
raise typer.Exit(1)
if unset:
if unset not in block.options:
console.print(f"[yellow]Option '{unset}' not set on {alias}[/yellow]")
return
del block.options[unset]
action = f"unset {unset}"
else:
block.set(option, value)
action = f"set {option} = {value}"
_replace_block_in_file(file_path, alias, block)
console.print(f"[green]Updated {alias}:[/green] {action}")
@host_app.command("remove")
def host_remove(
alias: Annotated[str, typer.Argument(help="Host alias to remove")],
force: Annotated[bool, typer.Option("--force", "-f", help="Skip confirmation")] = False,
) -> None:
"""Remove a host from its config file."""
from sshctl.operations.configops import find_host_file, remove_host_from_file
file_path = find_host_file(alias, _ssh_dir())
if not file_path:
console.print(f"[red]Host '{alias}' not found in any config file[/red]")
raise typer.Exit(1)
if not force:
from rich.prompt import Confirm
if not Confirm.ask(f"Remove host '{alias}' from {file_path}?"):
console.print("[yellow]Cancelled[/yellow]")
return
if remove_host_from_file(alias, file_path):
console.print(f"[green]Removed '{alias}' from {file_path}[/green]")
else:
console.print(f"[red]Failed to remove '{alias}'[/red]")
raise typer.Exit(1)
def _replace_block_in_file(file_path: Path, alias: str, block: HostBlock) -> None:
from sshctl.config.generator import generate_host_block
lines = file_path.read_text().splitlines(keepends=False)
result: list[str] = []
in_target = False
new_block = generate_host_block(block)
replaced = False
for line in lines:
stripped = line.strip()
if stripped.startswith("Host "):
aliases = stripped[4:].strip().split()
if alias in aliases:
in_target = True
if not replaced:
result.append(new_block)
replaced = True
continue
in_target = False
if in_target:
if not stripped or stripped.startswith("#"):
in_target = False
result.append(line)
continue
else:
result.append(line)
file_path.write_text("\n".join(result) + "\n")
jump_app = typer.Typer(help="Manage ProxyJump (enable/disable)", no_args_is_help=True)
app.add_typer(jump_app, name="jump")
@jump_app.command("off")
def jump_off() -> None:
"""Disable ProxyJump: save and remove all ProxyJump lines."""
from sshctl.operations.jump import off as jump_off_op
result = jump_off_op(_ssh_dir())
if result["status"] == "noop":
console.print("[yellow]" + str(result["message"]) + "[/yellow]")
return
count = result["count"]
nfiles = len(result["files"])
console.print(f"[green]Disabled {count} ProxyJump(s) across {nfiles} file(s)[/green]")
for f in result["files"]:
console.print(f" [dim]{f}[/dim]")
@jump_app.command("on")
def jump_on() -> None:
"""Enable ProxyJump: restore saved ProxyJump lines."""
from sshctl.operations.jump import on as jump_on_op
result = jump_on_op(_ssh_dir())
if result["status"] == "noop":
console.print(f"[yellow]{result['message']}[/yellow]")
return
count = result["count"]
nfiles = len(result["files"])
console.print(f"[green]Restored {count} ProxyJump(s) across {nfiles} file(s)[/green]")
for f in result["files"]:
console.print(f" [dim]{f}[/dim]")
@jump_app.command("status")
def jump_status() -> None:
"""Show hosts with ProxyJump and current state."""
from sshctl.operations.jump import status as jump_status_op
info = jump_status_op(_ssh_dir())
table = Table(title=f"ProxyJump Hosts ({info['count']})")
table.add_column("Host", style="cyan")
table.add_column("Jump Server")
table.add_column("File")
for h in info["hosts"]:
table.add_row(h["alias"], h["jump"], h["file"])
console.print(table)
if info["saved"]:
console.print("[yellow]ProxyJump state saved (use 'jump on' to restore)[/yellow]")
@jump_app.command("auto")
def jump_auto(
subnet_check: Annotated[
bool,
typer.Option(
"--subnet/--no-subnet",
help="Probe jump gateways for subnet reachability",
),
] = True,
) -> None:
"""Auto-toggle ProxyJump based on network environment.
Uses subnet probing (TCP to jump gateways) by default,
falls back to SSID matching.
"""
from sshctl.operations.jump import (
current_ssid,
is_internal_by_ssid,
is_internal_by_subnet,
known_networks,
)
from sshctl.operations.jump import (
off as jump_off_op,
)
from sshctl.operations.jump import (
on as jump_on_op,
)
internal = None
reasons: list[str] = []
if subnet_check:
reachable, probes = is_internal_by_subnet(_ssh_dir())
reachable_gws = [p for p in probes if p["reachable"]]
if reachable_gws:
internal = True
names = ", ".join(p["gateway"] for p in reachable_gws)
reasons.append(f"reachable via subnet: {names}")
else:
reasons.append("no jump gateway reachable on this subnet")
probes_table = Table(title="Jump Gateway Probes")
probes_table.add_column("Gateway", style="cyan")
probes_table.add_column("IP", style="dim")
probes_table.add_column("Port")
probes_table.add_column("Status")
probes_table.add_column("Detail")
for p in probes:
st = "green" if p["reachable"] else "red"
probes_table.add_row(
p["gateway"],
p["ip"] or "-",
str(p.get("port", 22)),
f"[{st}]{'reachable' if p['reachable'] else 'unreachable'}[/{st}]",
p["error"],
)
console.print(probes_table)
if internal is None:
ssid = current_ssid()
networks = known_networks(_ssh_dir())
reasons.append(f"SSID: {ssid or 'unknown'}")
reasons.append(f"patterns: {', '.join(networks)}")
if ssid:
internal = is_internal_by_ssid(_ssh_dir())
if internal is True:
reasons.append(f"SSID '{ssid}' matches internal patterns")
elif internal is False:
reasons.append(f"SSID '{ssid}' does not match internal patterns")
if internal is None:
console.print("[red]Could not determine network environment[/red]")
console.print(" - Install nmcli for SSID detection")
console.print(" - Or add SSID patterns to ~/.ssh/.internal_ssids")
raise typer.Exit(1)
color = "green" if internal else "yellow"
label = "internal" if internal else "external"
detail = "; ".join(reasons)
console.print(f"Decision: [{color}]{label}[/] — {detail}")
if internal:
result = jump_off_op(_ssh_dir())
if result["status"] == "noop":
console.print("[yellow]ProxyJump already off[/yellow]")
else:
count = result["count"]
nfiles = len(result["files"])
console.print(f"[green]Disabled {count} ProxyJump(s) across {nfiles} file(s)[/green]")
else:
result = jump_on_op(_ssh_dir())
if result["status"] == "noop":
console.print("[yellow]Already on (or no saved state)[/yellow]")
else:
count = result["count"]
nfiles = len(result["files"])
console.print(f"[green]Restored {count} ProxyJump(s) across {nfiles} file(s)[/green]")
@jump_app.command("networks")
def jump_networks(
action: Annotated[
str | None,
typer.Argument(help="list, add, remove"),
] = "list",
ssid: Annotated[
str | None,
typer.Option("--ssid", "-s", help="SSID or pattern (supports *)"),
] = None,
) -> None:
"""Manage internal network SSID list for auto-detection."""
from sshctl.operations.jump import (
add_network,
current_ssid,
known_networks,
)
path = _ssh_dir() / ".internal_ssids"
if action == "list":
nets = known_networks(_ssh_dir())
if not nets:
console.print("[yellow]No internal networks configured[/yellow]")
console.print("Add with: sshctl jump networks add --ssid <pattern>")
return
console.print("[bold]Internal Network Patterns:[/bold]")
for n in nets:
console.print(f" [cyan]{n}[/cyan]")
cur = current_ssid()
if cur:
matched = any(
re.match(p.replace("*", ".*").replace("?", ".") + "$", cur, re.I) for p in nets
)
if matched:
console.print(f"\nCurrent [green]{cur}[/green] is [bold]internal[/bold]")
else:
console.print(f"\nCurrent [yellow]{cur}[/yellow] is [bold]external[/bold]")
elif action == "add":
if not ssid:
ssid = current_ssid()
if not ssid:
console.print("[red]No SSID specified and could not detect current network[/red]")
raise typer.Exit(1)
add_network(_ssh_dir(), ssid)
console.print(f"[green]Added '{ssid}' to internal networks[/green]")
elif action == "remove":
if not ssid:
console.print("[red]Specify --ssid to remove[/red]")
raise typer.Exit(1)
if not path.exists():
console.print("[yellow]No internal networks file[/yellow]")
return
lines = path.read_text().splitlines()
new_lines = [ln for ln in lines if ln.strip() != ssid]
path.write_text("\n".join(new_lines) + "\n")
console.print(f"[green]Removed '{ssid}' from internal networks[/green]")
@app.command()
def check(
pattern: Annotated[
str | None,
typer.Option("--pattern", "-p", help="Glob pattern to filter hosts"),
] = None,
timeout: Annotated[
int, typer.Option("--timeout", "-t", help="Timeout per host (seconds)")
] = 10,
max_concurrent: Annotated[
int,
typer.Option("--max-concurrent", "-c", help="Max parallel checks"),
] = 50,
) -> None:
"""Check connectivity to all hosts."""
from sshctl.operations.checker import check_all_hosts
with console.status("[bold green]Checking hosts..."):
try:
results = asyncio.run(
check_all_hosts(
_ssh_config_path(),
timeout=timeout,
max_concurrent=max_concurrent,
filter_alias=pattern,
)
)
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
passed = sum(1 for r in results if r.passed)
failed = sum(1 for r in results if not r.passed)
table = Table(title=f"Check Results ({passed} passed, {failed} failed)")
table.add_column("Host", style="cyan")
table.add_column("HostName")
table.add_column("Port")
table.add_column("ProxyJump")
table.add_column("Status")
for r in results:
status_style = "green" if r.passed else "red"
table.add_row(
r.alias,
r.hostname,
str(r.port),
r.proxy_jump,
f"[{status_style}]{r.summary}[/{status_style}]",
)
console.print(table)
if failed > 0:
console.print(f"\n[yellow]{failed} host(s) failed check[/yellow]")
for r in results:
if not r.passed:
details = " | ".join(f for f in [r.dns_error, r.tcp_error, r.ssh_error] if f)
console.print(f" [red]{r.alias}:[/red] {details}")
raise typer.Exit(1)
@app.command()
def exec(
command: Annotated[str, typer.Argument(help="Command to run on each host")],
pattern: Annotated[
str | None,
typer.Option("--pattern", "-p", help="Filter hosts by glob pattern"),
] = None,
max_concurrent: Annotated[
int,
typer.Option("--max-concurrent", "-c", help="Max parallel executions"),
] = 20,
) -> None:
"""Run a command in parallel across hosts."""
from sshctl.operations.executor import run_on_all_hosts
with console.status(f"[bold green]Running '{command}' on hosts..."):
try:
results = asyncio.run(
run_on_all_hosts(
command,
_ssh_config_path(),
max_concurrent=max_concurrent,
filter_alias=pattern,
)
)
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
for r in results:
if r.error:
console.print(f"[red]{r.alias}:[/red] {r.error}")
elif r.returncode == 0:
if r.stdout:
console.print(f"[green]{r.alias}:[/green]\n{r.stdout}")
else:
console.print(f"[red]{r.alias} (exit {r.returncode}):[/red]")
if r.stderr:
console.print(r.stderr)
console.print(f"\n[bold]Executed on {len(results)} hosts[/bold]")
@app.command()
def csv(
output: Annotated[
Path | None,
typer.Option("--output", "-o", help="Output CSV file"),
] = None,
) -> None:
"""Export host inventory to CSV."""
from sshctl.export.exporter import export_csv_to_file
with console.status("[bold green]Exporting inventory..."):
try:
if output:
path = export_csv_to_file(output, config_path=_ssh_config_path())
console.print(f"[green]Exported to {path}[/green]")
else:
from sshctl.export.exporter import export_csv
content = export_csv(config_path=_ssh_config_path())
console.print(content)
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
@app.command()
def chmod(
ssh_dir: Annotated[
Path | None,
typer.Argument(help="SSH directory (default: ~/.ssh)"),
] = None,
check_only: Annotated[
bool,
typer.Option("--check", "-c", help="Only check, don't fix"),
] = False,
) -> None:
"""Fix or check SSH directory permissions."""
from sshctl.export.permissions import check_permissions, fix_ssh_permissions
target = ssh_dir.expanduser().resolve() if ssh_dir else _ssh_dir()
try:
if check_only:
issues = check_permissions(target)
total = sum(len(v) for v in issues.values())
if total == 0:
console.print("[green]All permissions correct[/green]")
else:
for category, paths in issues.items():
if paths:
console.print(f"[yellow]{category}: {len(paths)} issues[/yellow]")
for p in paths[:5]:
console.print(f" {p}")
raise typer.Exit(1)
else:
counts = fix_ssh_permissions(target)
table = Table(title=f"Permissions fixed in {target}")
table.add_column("Type", style="cyan")
table.add_column("Count")
for k, v in counts.items():
table.add_row(k, str(v))
console.print(table)
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
@app.command()
def ansible(
output_dir: Annotated[
Path | None,
typer.Option("--output", "-o", help="Output directory for inventory"),
] = None,
) -> None:
"""Generate Ansible inventory from SSH config."""
from sshctl.export.inventory import write_inventory
with console.status("[bold green]Generating Ansible inventory..."):
try:
path = write_inventory(_ssh_dir(), output_dir)
console.print(f"[green]Inventory written to {path}[/green]")
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
@app.command()
def validate(
fix: Annotated[bool, typer.Option("--fix", "-f", help="Auto-fix detected issues")] = False,
) -> None:
"""Validate SSH config structure and option correctness."""
from sshctl.config.validator import validate_config
from sshctl.operations.configops import heal_config
if fix:
actions = heal_config(_ssh_config_path(), auto_fix=True)
if actions:
table = Table(title="Config Healing Results")
table.add_column("Action", style="cyan")
table.add_column("Detail")
for a in actions:
st = {"removed": "red", "fixed": "green", "warning": "yellow", "detected": "yellow"}
table.add_row(f"[{st.get(a['action'], 'white')}]{a['action']}[/]", a["message"])
console.print(table)
try:
result = validate_config(_ssh_config_path())
except Exception as e:
console.print(f"[red]Error:[/red] {e}")
raise typer.Exit(1) from e
if result.passed:
console.print("[green]Config is valid[/green]")
else:
console.print(f"[red]Found {len(result.issues)} issues:[/red]")
for issue in result.issues:
sev_style = "red" if issue.severity == "error" else "yellow"
loc = f"{issue.file}:{issue.alias}" if issue.alias else issue.file
console.print(f" [{sev_style}]{issue.severity}[/{sev_style}] {loc}: {issue.message}")
raise typer.Exit(1)
@app.command()
def heal(
dry_run: Annotated[
bool, typer.Option("--dry-run", "-n", help="Show issues without fixing")
] = False,
) -> None:
"""Detect and fix SSH config issues (typos, invalid options)."""
from sshctl.operations.configops import heal_config
actions = heal_config(_ssh_config_path(), auto_fix=not dry_run)
if not actions:
console.print("[green]No issues found[/green]")
return
table = Table(title="Issues Found")
table.add_column("Status", style="cyan")
table.add_column("Detail")
for a in actions:
st = {"removed": "red", "fixed": "green", "warning": "yellow", "detected": "yellow"}
table.add_row(f"[{st.get(a['action'], 'white')}]{a['action']}[/]", a["message"])
console.print(table)
if dry_run:
n = len(actions)
console.print(f"\n[yellow]{n} issue(s) detected. Run without --dry-run to fix.[/yellow]")
@app.command()
def log(
lines: Annotated[int, typer.Option("--lines", "-n", help="Number of recent events")] = 50,
follow: Annotated[bool, typer.Option("--follow", "-f", help="Watch for new events")] = False,
) -> None:
"""Show recent sshctl activity log."""
from sshctl.operations.events import recent
events = recent(_ssh_dir(), limit=lines)
if not events:
console.print("[yellow]No events logged yet[/yellow]")
return
table = Table(title=f"Recent Events ({len(events)})")
table.add_column("Time", style="dim")
table.add_column("Command", style="cyan")
table.add_column("Status")
table.add_column("Detail")
for ev in reversed(events):
from datetime import datetime
ts = datetime.fromtimestamp(ev.get("ts", 0)).strftime("%H:%M:%S")
cmd = ev.get("cmd", "")
status = ev.get("status", "")
detail = ev.get("result", "")
st = "green" if status == "ok" else "red" if status == "error" else "yellow"
table.add_row(ts, cmd, f"[{st}]{status}[/{st}]", detail)
console.print(table)
if follow:
import time
try:
last_count = len(events)
while True:
time.sleep(2)
new_events = recent(_ssh_dir(), limit=lines)
if len(new_events) > last_count:
for ev in new_events[last_count - len(new_events) :]:
ts = datetime.fromtimestamp(ev.get("ts", 0)).strftime("%H:%M:%S")
status = ev.get("status", "")
cmd = ev.get("cmd", "")
result = ev.get("result", "")
console.print(f" {ts} [{status}] {cmd}: {result}")
last_count = len(new_events)
except KeyboardInterrupt:
pass
@app.command()
def diff() -> None:
"""Show differences between generated and current SSH config."""
from difflib import unified_diff
from sshctl.config.generator import generate_config_text
from sshctl.config.parser import resolve_includes
config_path = _ssh_config_path()
if not config_path.exists():
console.print("[red]Config file not found[/red]")
raise typer.Exit(1)
try:
config = resolve_includes(config_path)
current = config_path.read_text(encoding="utf-8")
except Exception as e:
console.print(f"[red]Error reading config:[/red] {e}")
raise typer.Exit(1) from e
generated = generate_config_text(config)
if current == generated:
console.print("[green]Config is up to date, no differences[/green]")
return
diff_lines = list(
unified_diff(
current.splitlines(keepends=True),
generated.splitlines(keepends=True),
fromfile=str(config_path),
tofile="generated",
)
)
for line in diff_lines:
if line.startswith("+"):
console.print(f"[green]{line.rstrip()}[/green]")
elif line.startswith("-"):
console.print(f"[red]{line.rstrip()}[/red]")
elif line.startswith("@@"):
console.print(f"[cyan]{line.rstrip()}[/cyan]")
else:
console.print(line.rstrip())
@app.command()
def backup(
output_dir: Annotated[
Path | None,
typer.Option("--output", "-o", help="Backup output directory"),
] = None,
) -> None:
"""Create a timestamped backup of the SSH config directory."""
import shutil
from datetime import datetime
ssh_dir = _ssh_dir()
if not ssh_dir.exists():
console.print("[red]SSH directory not found[/red]")
raise typer.Exit(1)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
backup_name = f"ssh_backup_{timestamp}"
output_dir = ssh_dir.parent / backup_name if output_dir is None else output_dir / backup_name
output_dir.mkdir(parents=True, exist_ok=True)
try:
shutil.copytree(
ssh_dir, output_dir, dirs_exist_ok=True, ignore=shutil.ignore_patterns("*.sock")
)
console.print(f"[green]Backup created at {output_dir}[/green]")
except Exception as e:
console.print(f"[red]Backup failed:[/red] {e}")
raise typer.Exit(1) from e
@app.command()
def restore(
backup_path: Annotated[Path, typer.Argument(help="Path to backup directory")],
) -> None:
"""Restore SSH config from a backup."""
import shutil
if not backup_path.exists():
console.print(f"[red]Backup not found: {backup_path}[/red]")
raise typer.Exit(1)
ssh_dir = _ssh_dir()
if ssh_dir.exists():
import datetime
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
pre_restore = ssh_dir.parent / f"ssh_pre_restore_{timestamp}"
shutil.copytree(ssh_dir, pre_restore, dirs_exist_ok=True)
console.print(f"[yellow]Pre-restore backup saved at {pre_restore}[/yellow]")
shutil.rmtree(ssh_dir)
try:
shutil.copytree(backup_path, ssh_dir, dirs_exist_ok=True)
console.print(f"[green]Restored from {backup_path}[/green]")
except Exception as e:
console.print(f"[red]Restore failed:[/red] {e}")
raise typer.Exit(1) from e
@app.command()
def tree() -> None:
"""Show the conf.d directory tree."""
conf_d = _ssh_dir() / "conf.d"
if not conf_d.exists():
console.print("[yellow]No conf.d/ directory found[/yellow]")
raise typer.Exit(1)
_tree = Tree(f":file_folder: {conf_d}")
_walk_tree(conf_d, _tree)
console.print(_tree)
def _walk_tree(path: Path, node: Tree) -> None:
for entry in sorted(path.iterdir()):
if entry.name.startswith("."):
continue
if entry.is_dir():
branch = node.add(f":file_folder: [cyan]{entry.name}/[/cyan]")
_walk_tree(entry, branch)
elif entry.suffix == ".conf":
from sshctl.config.parser import parse_file
try:
parsed = parse_file(entry)
aliases = [b.alias for b in parsed.blocks if b.alias != "*"]
label = f"[green]{entry.name}[/green]"
if aliases:
label += f" ({len(aliases)} hosts)"
node.add(label)
except Exception:
node.add(f"[dim]{entry.name}[/dim]")
@app.command()
def cleanup(
dry_run: Annotated[
bool, typer.Option("--dry-run", "-n", help="Show stale sockets without removing")
] = False,
) -> None:
"""Remove stale SSH control master sockets."""
from sshctl.operations.cleanup import cleanup_control_sockets
result = cleanup_control_sockets(_ssh_dir(), dry_run=dry_run)
if result["total"] == 0:
console.print("[green]No control sockets found[/green]")
return
total = result["total"]
active = result["active"]
removed = result["removed"]
console.print(f"Total: {total}, Active: {active}, Removed: {removed}")
if dry_run and result["removed"]:
console.print(f"[yellow]{result['removed']} socket(s) would be removed[/yellow]")
_list_stale_sockets(_ssh_dir())
elif result["removed"]:
console.print(f"[green]Removed {result['removed']} stale socket(s)[/green]")
def _list_stale_sockets(ssh_dir: Path) -> None:
from sshctl.operations.cleanup import _is_socket_in_use
for sock in sorted((ssh_dir / "control.d").glob("*.sock")):
if not _is_socket_in_use(sock):
console.print(f" [dim]{sock.name}[/dim]")
setup_app = typer.Typer(help="Setup and manage SSH infrastructure", no_args_is_help=True)
app.add_typer(setup_app, name="setup")
@setup_app.command("completions")
def setup_completions(
shell: Annotated[
str | None,
typer.Option("--shell", "-s", help="Shell type (bash, zsh, fish)"),
] = None,
) -> None:
"""Install shell completions for sshctl."""
if shell:
_shell = shell.lower()
else:
import os
_shell = os.environ.get("SHELL", "").split("/")[-1] or "bash"
_paths = {
"bash": Path("~/.bash_completion.d/sshctl").expanduser(),
"zsh": Path("~/.zsh/completions/_sshctl").expanduser(),
"fish": Path("~/.config/fish/completions/sshctl.fish").expanduser(),
}
if _shell not in _paths:
console.print(f"[red]Unsupported shell: {_shell}[/red]")
raise typer.Exit(1)
from click.shell_completion import BashComplete, FishComplete, ZshComplete
_cls = {"bash": BashComplete, "zsh": ZshComplete, "fish": FishComplete}[_shell]
out = _paths[_shell]
out.parent.mkdir(parents=True, exist_ok=True)
cli = typer.main.get_command(app)
src = _cls(
cli=cast(click.Command, cli),
ctx_args={},
prog_name="sshctl",
complete_var="_SSHCTL_COMPLETE",
).source()
out.write_text(src)
out.chmod(0o644)
console.print(f"[green]Completions installed for {_shell} at {out}[/green]")
hints = {
"bash": f"echo 'source {out}' >> ~/.bashrc",
"zsh": f"echo 'fpath=({out.parent.parent} $fpath)' >> ~/.zshrc",
"fish": "",
}
if _shell in hints and hints[_shell]:
console.print(f"[yellow]Add to your .{_shell}rc:[/yellow]\n {hints[_shell]}")
@setup_app.command("agent")
def setup_agent(
action: Annotated[
str,
typer.Argument(help="Action: status, start, stop, restart"),
] = "status",
) -> None:
"""Manage the ssh-agent systemd service."""
from sshctl.operations.setup import agent_start, agent_status, agent_stop
match action:
case "start" | "restart":
result = agent_start()
case "stop":
result = agent_stop()
case _:
result = agent_status()
if "error" in result:
console.print(f"[red]{result['error']}[/red]")
raise typer.Exit(1)
console.print(f"[green]{result['service']}[/green] is [bold]{result['status']}[/bold]")
@setup_app.command("timer")
def setup_timer(
name: Annotated[
str,
typer.Argument(help="Timer: cleanup, check"),
] = "cleanup",
action: Annotated[
str,
typer.Option("--action", "-a", help="enable, disable, status"),
] = "status",
) -> None:
"""Manage systemd timers for periodic tasks."""
from sshctl.operations.setup import timer_disable, timer_enable, timer_status
if action == "enable":
result = timer_enable(name)
elif action == "disable":
result = timer_disable(name)
else:
results = timer_status()
table = Table(title="Systemd Timers")
table.add_column("Timer", style="cyan")
table.add_column("Description")
table.add_column("Status")
for r in results:
style = "green" if r["status"] == "active" else "yellow"
table.add_row(r["timer"], r["description"], f"[{style}]{r['status']}[/{style}]")
console.print(table)
return
if "error" in result:
console.print(f"[red]{result['error']}[/red]")
raise typer.Exit(1)
style = "green" if result.get("status") == "active" else "yellow"
console.print(f"Timer [cyan]{name}[/cyan]: [{style}]{result['status']}[/{style}]")
@setup_app.command("git")
def setup_git() -> None:
"""Initialize git tracking for SSH config files."""
from sshctl.operations.setup import init_git
result = init_git(_ssh_dir())
if result["status"] == "created":
console.print(f"[green]Git repo initialized at {result['path']}[/green]")
elif result["status"] == "exists":
console.print(f"[yellow]Git repo already exists at {result['path']}[/yellow]")
else:
console.print(f"[red]{result.get('message', 'Failed')}[/red]")
raise typer.Exit(1)
@setup_app.command("init")
def setup_init() -> None:
"""Full setup: git, agent, timers, completions."""
from sshctl.operations.setup import agent_start, init_git, timer_enable
table = Table(title="SSH Setup Results")
table.add_column("Component", style="cyan")
table.add_column("Status", style="green")
git_result = init_git(_ssh_dir())
table.add_row("git", str(git_result["status"]))
agent_result = agent_start()
table.add_row("ssh-agent", str(agent_result.get("status", "error")))
for t in ("cleanup", "check"):
timer_result = timer_enable(t)
table.add_row(f"timer/{t}", str(timer_result.get("status", "error")))
import os
shell = os.environ.get("SHELL", "").split("/")[-1] or "bash"
if shell in ("bash", "zsh", "fish"):
_install_completions(shell)
table.add_row(f"completions/{shell}", "installed")
console.print(table)
console.print("\n[yellow]Add to your shell rc:[/yellow]")
console.print(' export SSH_AUTH_SOCK="${XDG_RUNTIME_DIR}/ssh-agent.socket"')
def _install_completions(shell: str) -> None:
from click.shell_completion import BashComplete, FishComplete, ZshComplete
_paths = {
"bash": Path("~/.bash_completion.d/sshctl").expanduser(),
"zsh": Path("~/.zsh/completions/_sshctl").expanduser(),
"fish": Path("~/.config/fish/completions/sshctl.fish").expanduser(),
}
_cls = {"bash": BashComplete, "zsh": ZshComplete, "fish": FishComplete}
out = _paths[shell]
out.parent.mkdir(parents=True, exist_ok=True)
cli = typer.main.get_command(app)
src = _cls[shell](
cli=cast(click.Command, cli),
ctx_args={},
prog_name="sshctl",
complete_var="_SSHCTL_COMPLETE",
).source()
out.write_text(src)
out.chmod(0o644)
def serve() -> None:
from sshctl.operations import events
cmds, kwargs = _parse_argv(sys.argv[1:])
cmd = " ".join(cmds) if cmds else ""
start = time.time()
ok = True
try:
app()
except SystemExit as e:
ok = e.code == 0 or e.code is None
raise
except Exception:
ok = False
raise
finally:
dur = (time.time() - start) * 1000
events.log_event(
events.Event(
command=cmd or "root",
status="ok" if ok else "error",
args=kwargs,
duration_ms=dur,
)
)
events.flush(_ssh_dir())
def _parse_argv(argv: list[str]) -> tuple[list[str], dict[str, str]]:
cmds: list[str] = []
kwargs: dict[str, str] = {}
for a in argv:
if a.startswith("--"):
break
cmds.append(a)
i = len(cmds)
while i < len(argv):
if argv[i].startswith("--"):
key = argv[i].lstrip("-")
if i + 1 < len(argv) and not argv[i + 1].startswith("-"):
kwargs[key] = argv[i + 1]
i += 2
else:
kwargs[key] = "true"
i += 1
else:
kwargs[argv[i]] = "true"
i += 1
return cmds, kwargs
if __name__ == "__main__":
serve()