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", ""]