From 60ba7e828652b40f4d2ee73e4cce2f0732e75eb3 Mon Sep 17 00:00:00 2001 From: Amir Husayn Panahifar Date: Sat, 12 Sep 2026 13:47:31 +0330 Subject: [PATCH] feat(config): add ssh config parser, resolver, validator, generator --- src/sshctl/config/__init__.py | 0 src/sshctl/config/generator.py | 71 ++++++++++++++++++ src/sshctl/config/parser.py | 129 +++++++++++++++++++++++++++++++++ src/sshctl/config/resolver.py | 85 ++++++++++++++++++++++ src/sshctl/config/validator.py | 89 +++++++++++++++++++++++ 5 files changed, 374 insertions(+) create mode 100644 src/sshctl/config/__init__.py create mode 100644 src/sshctl/config/generator.py create mode 100644 src/sshctl/config/parser.py create mode 100644 src/sshctl/config/resolver.py create mode 100644 src/sshctl/config/validator.py diff --git a/src/sshctl/config/__init__.py b/src/sshctl/config/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/sshctl/config/generator.py b/src/sshctl/config/generator.py new file mode 100644 index 0000000..3d20a83 --- /dev/null +++ b/src/sshctl/config/generator.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +from collections.abc import Sequence +from pathlib import Path + +from sshctl.core.models import HostBlock, SshConfig + + +def generate_config_text(config: SshConfig, *, include_header: bool = True) -> str: + lines: list[str] = [] + + if include_header: + lines.append("# Generated by sshctl — do not hand-edit") + lines.append("") + + for block in config.blocks: + lines.append(f"Host {block.alias}") + for key in _option_order(block): + val = block.get(key) + if val is not None: + lines.append(f" {key} {val}") + lines.append("") + + for inc in config.raw_includes: + lines.append(f"Include {inc}") + + return "\n".join(lines) + + +def write_config_file(config: SshConfig, path: str | Path) -> Path: + path = Path(path) + text = generate_config_text(config) + path.write_text(text, encoding="utf-8") + return path + + +def generate_host_block(host: HostBlock) -> str: + lines: list[str] = [f"Host {host.alias}"] + for key in _option_order(host): + val = host.get(key) + if val is not None: + lines.append(f" {key} {val}") + return "\n".join(lines) + + +def write_host_file(hosts: Sequence[HostBlock], path: str | Path) -> Path: + path = Path(path) + lines: list[str] = [f"# Generated by sshctl — {path.name}", ""] + for host in hosts: + lines.append(generate_host_block(host)) + lines.append("") + path.write_text("\n".join(lines), encoding="utf-8") + return path + + +_ORDERED_KEYS = [ + "HostName", + "User", + "Port", + "ProxyJump", + "IdentityFile", + "CertificateFile", + "IdentitiesOnly", + "PreferredAuthentications", +] + + +def _option_order(block: HostBlock) -> list[str]: + ordered = [k for k in _ORDERED_KEYS if k in block.options] + remaining = sorted(k for k in block.options if k not in _ORDERED_KEYS) + return ordered + remaining diff --git a/src/sshctl/config/parser.py b/src/sshctl/config/parser.py new file mode 100644 index 0000000..2b975e9 --- /dev/null +++ b/src/sshctl/config/parser.py @@ -0,0 +1,129 @@ +from __future__ import annotations + +import re +from collections import deque +from pathlib import Path + +from sshctl.core.exceptions import ParseError +from sshctl.core.models import HostBlock, SshConfig + +_RE_HOST = re.compile(r"^Host\s+(.+)$", re.I) +_RE_OPTION = re.compile(r"^\s+([a-zA-Z0-9]+)\s+(.+)$") +_RE_INCLUDE = re.compile(r"^Include\s+(.+)$", re.I) +_RE_COMMENT = re.compile(r"^\s*(#|$)") +_RE_CONTINUATION = re.compile(r"\s*\\\s*$") + + +def parse_config(text: str) -> SshConfig: + config = SshConfig() + current_block: HostBlock | None = None + lines = text.splitlines() + i = 0 + + while i < len(lines): + raw_line = lines[i] + line = raw_line.strip() + + if not line or _RE_COMMENT.match(line): + i += 1 + continue + + if _RE_CONTINUATION.match(line): + continued = raw_line.rstrip("\\").rstrip() + i += 1 + while i < len(lines): + next_line = lines[i].strip() + if _RE_CONTINUATION.match(lines[i]): + continued += " " + lines[i].strip().rstrip("\\").rstrip() + i += 1 + else: + continued += " " + next_line + i += 1 + break + line = continued + + host_match = _RE_HOST.match(line) + if host_match: + current_block = HostBlock(alias=host_match.group(1).strip()) + config.blocks.append(current_block) + i += 1 + continue + + include_match = _RE_INCLUDE.match(line) + if include_match: + config.raw_includes.append(include_match.group(1).strip()) + i += 1 + continue + + opt_match = _RE_OPTION.match(raw_line) + if opt_match and current_block: + key, value = opt_match.group(1), opt_match.group(2).strip() + value = value.strip('"').strip("'") + current_block.set(key, value) + i += 1 + continue + + if opt_match and not current_block: + i += 1 + continue + + i += 1 + + return config + + +def parse_file(path: str | Path) -> SshConfig: + path = Path(path) + if not path.exists(): + raise ParseError(f"Config file not found: {path}") + try: + text = path.read_text(encoding="utf-8") + except OSError as e: + raise ParseError(f"Failed to read config: {path}: {e}") from e + return parse_config(text) + + +def resolve_includes( + config_path: str | Path, + base_dir: str | Path | None = None, +) -> SshConfig: + config_path = Path(config_path) + if base_dir is None: + base_dir = config_path.parent + base_dir = Path(base_dir) + + merged = parse_file(config_path) + seen: set[Path] = {config_path.resolve()} + + to_process: deque[str] = deque(merged.raw_includes) + _resolve_include_chain(merged, to_process, base_dir, seen) + + return merged + + +def _resolve_include_chain( + merged: SshConfig, + patterns: deque[str], + base_dir: Path, + seen: set[Path], +) -> None: + while patterns: + pattern = patterns.popleft() + if "*" in pattern or "?" in pattern: + matched = sorted(base_dir.glob(pattern)) + else: + matched = [base_dir / pattern] + + for matched_path in matched: + matched_path = matched_path.resolve() + if matched_path in seen or not matched_path.exists() or matched_path.is_dir(): + continue + seen.add(matched_path) + try: + included = parse_file(matched_path) + merged.blocks.extend(included.blocks) + for inc in included.raw_includes: + if inc not in patterns: + patterns.append(inc) + except ParseError: + pass diff --git a/src/sshctl/config/resolver.py b/src/sshctl/config/resolver.py new file mode 100644 index 0000000..14cddd0 --- /dev/null +++ b/src/sshctl/config/resolver.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +import shlex +import subprocess +from pathlib import Path +from typing import NamedTuple + +from sshctl.core.exceptions import ResolverError + + +class ResolvedHost(NamedTuple): + alias: str + hostname: str + user: str + port: int + proxy_jump: str + identity_file: str + options: dict[str, str] + + +def resolve_host( + alias: str, + config_path: str | Path, + *, + ssh_bin: str = "ssh", +) -> ResolvedHost: + cmd = [ssh_bin, "-F", str(config_path), "-G", alias] + try: + result = subprocess.run( + cmd, + capture_output=True, + text=True, + timeout=30, + ) + except (subprocess.TimeoutExpired, FileNotFoundError) as e: + raise ResolverError(f"Failed to resolve {alias}: {e}") from e + + if result.returncode != 0: + raise ResolverError( + f"Failed to resolve {alias}: {result.stderr.strip() or 'unknown error'}" + ) + + opts = _parse_ssh_g_output(result.stdout) + return ResolvedHost( + alias=alias, + hostname=opts.get("hostname", alias), + user=opts.get("user", "root"), + port=int(opts.get("port", "22")), + proxy_jump=opts.get("proxyjump", ""), + identity_file=opts.get("identityfile", ""), + options=opts, + ) + + +def resolve_all_hosts( + config_path: str | Path, + *, + ssh_bin: str = "ssh", +) -> list[ResolvedHost]: + from sshctl.config.parser import resolve_includes + + config = resolve_includes(config_path) + results: list[ResolvedHost] = [] + for alias in config.aliases: + try: + resolved = resolve_host(alias, config_path, ssh_bin=ssh_bin) + results.append(resolved) + except ResolverError: + continue + return results + + +def _parse_ssh_g_output(text: str) -> dict[str, str]: + opts: dict[str, str] = {} + for line in text.splitlines(): + line = line.strip() + if not line: + continue + parts = shlex.split(line) + if len(parts) >= 2: + key = parts[0].lower() + val = " ".join(parts[1:]).strip('"').strip("'") + if key not in opts: + opts[key] = val + return opts diff --git a/src/sshctl/config/validator.py b/src/sshctl/config/validator.py new file mode 100644 index 0000000..6fd5ca1 --- /dev/null +++ b/src/sshctl/config/validator.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path + +from sshctl.config.parser import parse_file, resolve_includes +from sshctl.core.models import HostBlock + + +@dataclass +class ValidationIssue: + severity: str + message: str + file: str = "" + alias: str = "" + + +@dataclass +class ValidationResult: + passed: bool = True + issues: list[ValidationIssue] = field(default_factory=list) + + def add(self, severity: str, message: str, file: str = "", alias: str = "") -> None: + self.issues.append(ValidationIssue(severity, message, file, alias)) + if severity in ("error",): + self.passed = False + + +def validate_config(config_path: str | Path) -> ValidationResult: + config_path = Path(config_path) + result = ValidationResult() + + if not config_path.exists(): + result.add("error", f"Config file not found: {config_path}") + return result + + try: + config = resolve_includes(config_path) + except Exception as e: + result.add("error", f"Failed to parse config: {e}") + return result + + for block in config.blocks: + _validate_host_block(block, config_path, result) + + _validate_includes(config_path, result) + + return result + + +def _validate_host_block( + block: HostBlock, + config_path: Path, + result: ValidationResult, +) -> None: + if not block.alias: + result.add("error", "Host block with empty alias", str(config_path)) + return + + if block.alias == "*": + return + + if not block.get("HostName"): + result.add( + "warning", + f"Host '{block.alias}' has no HostName", + str(config_path), + block.alias, + ) + + +def _validate_includes(config_path: Path, result: ValidationResult) -> None: + try: + raw = parse_file(config_path) + except Exception: + return + + base = config_path.parent + for inc_pattern in raw.raw_includes: + if "*" in inc_pattern or "?" in inc_pattern: + matched = list(base.glob(inc_pattern)) + if not matched: + result.add("warning", f"Include pattern '{inc_pattern}' matches no files") + else: + target = base / inc_pattern + if target.suffix != ".conf" and target.is_dir(): + continue + if not target.exists(): + result.add("warning", f"Include target '{inc_pattern}' not found")