Files

294 lines
9.4 KiB
Python
Raw Permalink Normal View History

from __future__ import annotations
from pathlib import Path
from sshctl.operations.configops import (
_clean_blank_lines,
find_host_file,
get_host_aliases,
get_host_star_block,
heal_config,
remove_host_from_file,
set_host_star_option,
unset_host_star_option,
validate_options,
)
STAR_CONFIG = """Host *
User root
Port 22
StrictHostKeyChecking accept-new
Include conf.d/
"""
def _write(path: Path, text: str) -> Path:
path.write_text(text)
return path
def test_get_host_star_block(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", STAR_CONFIG)
opts = get_host_star_block(cfg)
assert opts["User"] == "root"
assert opts["Port"] == "22"
assert opts["StrictHostKeyChecking"] == "accept-new"
def test_get_host_star_block_no_star(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host web\n HostName 10.0.0.1\n")
assert get_host_star_block(cfg) == {}
def test_get_host_star_block_comments_skipped(tmp_path: Path) -> None:
cfg = _write(
tmp_path / "config",
"Host *\n ## comment\n User root\n # another\n Port 22\n",
)
opts = get_host_star_block(cfg)
assert opts == {"User": "root", "Port": "22"}
def test_get_host_star_block_stops_at_include(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", STAR_CONFIG)
opts = get_host_star_block(cfg)
assert "after_include" not in opts
def test_get_host_star_block_stripped_quotes(tmp_path: Path) -> None:
cfg = _write(
tmp_path / "config",
"Host *\n IdentityFile \"~/.ssh/id_ed25519\"\n ProxyJump 'gw'\n",
)
opts = get_host_star_block(cfg)
assert opts["IdentityFile"] == "~/.ssh/id_ed25519"
assert opts["ProxyJump"] == "gw"
def test_set_host_star_option_update_existing(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", STAR_CONFIG)
set_host_star_option(cfg, "Port", "2222")
opts = get_host_star_block(cfg)
assert opts["Port"] == "2222"
text = cfg.read_text()
assert "Port 2222" in text
def test_set_host_star_option_add_new(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n User root\n\nInclude conf.d/\n")
set_host_star_option(cfg, "Port", "2222")
opts = get_host_star_block(cfg)
assert opts["Port"] == "2222"
def test_set_host_star_option_case_insensitive(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n port 22\n")
set_host_star_option(cfg, "Port", "2222")
opts = get_host_star_block(cfg)
assert opts["Port"] == "2222"
def test_set_host_star_option_no_star_block(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host web\n HostName 10.0.0.1\n")
set_host_star_option(cfg, "Port", "22")
assert "Port 22" in cfg.read_text()
def test_unset_host_star_option(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", STAR_CONFIG)
assert unset_host_star_option(cfg, "Port") is True
assert "Port" not in get_host_star_block(cfg)
assert "Port 22" not in cfg.read_text()
def test_unset_host_star_option_missing(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n User root\n")
assert unset_host_star_option(cfg, "Port") is False
def test_unset_host_star_option_case_insensitive(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n port 22\n")
assert unset_host_star_option(cfg, "Port") is True
def test_validate_options_valid(tmp_path: Path) -> None:
cfg = _write(
tmp_path / "config",
"Host *\n User root\n Port 22\n StrictHostKeyChecking accept-new\n",
)
assert validate_options(cfg) == []
def test_validate_options_unknown(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n FooBarBaz yes\n")
issues = validate_options(cfg)
assert len(issues) == 1
assert issues[0]["type"] == "unknown"
assert issues[0]["key"] == "FooBarBaz"
def test_validate_options_typo(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n hostname example.com\n")
issues = validate_options(cfg)
assert len(issues) == 1
assert issues[0]["type"] == "typo"
assert issues[0]["fixed"] == "HostName"
def test_validate_options_invalid(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n GSSAPIKeyExchange yes\n")
issues = validate_options(cfg)
assert len(issues) == 1
assert issues[0]["type"] == "invalid"
def test_validate_options_skips_comments_hosts_includes(tmp_path: Path) -> None:
cfg = _write(
tmp_path / "config",
"# comment\nHost *\n User root\nInclude conf.d/\n",
)
assert validate_options(cfg) == []
def test_validate_options_multiword_typo(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n strict hostkeychecking yes\n")
issues = validate_options(cfg)
assert len(issues) == 1
assert issues[0]["type"] == "unknown"
def test_validate_proper_casing_not_flagged(tmp_path: Path) -> None:
cfg = _write(
tmp_path / "config",
"Host *\n HostName example.com\n User root\n Port 22\n",
)
assert validate_options(cfg) == []
def test_heal_config_dry_run(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n hostname example.com\n")
actions = heal_config(cfg, auto_fix=False)
assert len(actions) == 1
assert actions[0]["action"] == "detected"
def test_heal_config_no_issues(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n User root\n")
assert heal_config(cfg, auto_fix=True) == []
def test_heal_config_fix_typo(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host web\n hostname example.com\n")
actions = heal_config(cfg, auto_fix=True)
assert any(a["action"] == "fixed" for a in actions)
assert "HostName example.com" in cfg.read_text()
def test_heal_config_keep_correct_case(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host web\n HostName example.com\n")
actions = heal_config(cfg, auto_fix=True)
assert all(a["action"] != "fixed" for a in actions)
def test_heal_config_remove_invalid(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n GSSAPIKeyExchange yes\n")
actions = heal_config(cfg, auto_fix=True)
assert any(a["action"] == "removed" for a in actions)
assert "GSSAPIKeyExchange" not in cfg.read_text()
def test_heal_config_warnings_remaining(tmp_path: Path) -> None:
cfg = _write(tmp_path / "config", "Host *\n FooBarBaz yes\n")
actions = heal_config(cfg, auto_fix=True)
assert any(a["action"] == "warning" for a in actions)
def test_find_host_file(ssh_dir: Path) -> None:
conf_d = ssh_dir / "conf.d"
conf_d.mkdir()
(conf_d / "10-servers.conf").write_text(
"Host web01\n HostName 10.0.0.1\n\nHost web02\n HostName 10.0.0.2\n"
)
found = find_host_file("web02", ssh_dir)
assert found is not None
assert found.name == "10-servers.conf"
def test_find_host_file_missing(ssh_dir: Path) -> None:
conf_d = ssh_dir / "conf.d"
conf_d.mkdir()
(conf_d / "10-servers.conf").write_text("Host web01\n HostName 10.0.0.1\n")
assert find_host_file("nope", ssh_dir) is None
def test_find_host_file_no_conf_d(ssh_dir: Path) -> None:
assert find_host_file("web01", ssh_dir) is None
def test_find_host_file_multi_alias_line(ssh_dir: Path) -> None:
conf_d = ssh_dir / "conf.d"
conf_d.mkdir()
(conf_d / "10-servers.conf").write_text("Host web01 web02\n HostName 10.0.0.1\n")
assert find_host_file("web02", ssh_dir) is not None
def test_get_host_aliases(ssh_dir: Path) -> None:
conf_d = ssh_dir / "conf.d"
conf_d.mkdir()
(conf_d / "10-servers.conf").write_text(
"Host *\n User root\nHost web01\n HostName 10.0.0.1\nHost web02\n HostName 10.0.0.2\n"
)
assert get_host_aliases(ssh_dir) == ["web01", "web02"]
def test_get_host_aliases_no_conf_d(ssh_dir: Path) -> None:
assert get_host_aliases(ssh_dir) == []
def test_get_host_aliases_recursive(ssh_dir: Path) -> None:
nested = ssh_dir / "conf.d" / "internal" / "prod"
nested.mkdir(parents=True)
(nested / "01-web.conf").write_text("Host app01\n HostName 10.0.0.1\n")
assert get_host_aliases(ssh_dir) == ["app01"]
def test_remove_host_from_file(tmp_path: Path) -> None:
f = _write(
tmp_path / "hosts.conf",
"Host web01\n HostName 10.0.0.1\n\nHost web02\n HostName 10.0.0.2\n",
)
assert remove_host_from_file("web01", f) is True
text = f.read_text()
assert "web01" not in text
assert "web02" in text
def test_remove_host_from_file_missing(tmp_path: Path) -> None:
f = _write(tmp_path / "hosts.conf", "Host web01\n HostName 10.0.0.1\n")
assert remove_host_from_file("nope", f) is False
def test_remove_host_cleans_duplicate_blanks(tmp_path: Path) -> None:
f = _write(
tmp_path / "hosts.conf",
"Host web01\n HostName 10.0.0.1\n\n\n\n\nHost web02\n HostName 10.0.0.2\n",
)
remove_host_from_file("web01", f)
text = f.read_text()
assert "\n\n\n" not in text
def test_remove_host_keeps_comments_before_next(tmp_path: Path) -> None:
f = _write(
tmp_path / "hosts.conf",
"Host web01\n HostName 10.0.0.1\n\n# next host\nHost web02\n HostName 10.0.0.2\n",
)
remove_host_from_file("web01", f)
assert "web02" in f.read_text()
def test_clean_blank_lines() -> None:
cleaned = _clean_blank_lines(["Host a", "", "", " HostName x", ""])
assert cleaned == ["Host a", "", " HostName x", ""]