feat(config): add ssh config parser, resolver, validator, generator
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user