diff --git a/src/sshctl/__init__.py b/src/sshctl/__init__.py new file mode 100644 index 0000000..b94ac3c --- /dev/null +++ b/src/sshctl/__init__.py @@ -0,0 +1,5 @@ +from __future__ import annotations + +from sshctl.core import __version__ + +__all__ = ["__version__"] diff --git a/src/sshctl/__main__.py b/src/sshctl/__main__.py new file mode 100644 index 0000000..5b11448 --- /dev/null +++ b/src/sshctl/__main__.py @@ -0,0 +1,5 @@ +from __future__ import annotations + +from sshctl.cli.app import app + +app() diff --git a/src/sshctl/cli/__init__.py b/src/sshctl/cli/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/sshctl/cli/app.py b/src/sshctl/cli/app.py new file mode 100644 index 0000000..e436636 --- /dev/null +++ b/src/sshctl/cli/app.py @@ -0,0 +1,1385 @@ +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 ") + 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()