From e3a01c4d3ab912ee075623aadd8fdd60a9cf2b77 Mon Sep 17 00:00:00 2001 From: Amir Husayn Panahifar Date: Sat, 12 Sep 2026 13:47:31 +0330 Subject: [PATCH] test: add unit tests for core, parser and operations --- tests/conftest.py | 86 +++++++++++ tests/test_checker.py | 119 ++++++++++++++++ tests/test_cleanup.py | 91 ++++++++++++ tests/test_config.py | 32 +++++ tests/test_configops.py | 293 ++++++++++++++++++++++++++++++++++++++ tests/test_events.py | 64 +++++++++ tests/test_exceptions.py | 25 ++++ tests/test_executor.py | 24 ++++ tests/test_exporter.py | 67 +++++++++ tests/test_generator.py | 103 ++++++++++++++ tests/test_inventory.py | 31 ++++ tests/test_jump.py | 271 +++++++++++++++++++++++++++++++++++ tests/test_logutil.py | 24 ++++ tests/test_models.py | 128 +++++++++++++++++ tests/test_parser.py | 136 ++++++++++++++++++ tests/test_permissions.py | 35 +++++ tests/test_resolver.py | 101 +++++++++++++ tests/test_validator.py | 21 +++ 18 files changed, 1651 insertions(+) create mode 100644 tests/conftest.py create mode 100644 tests/test_checker.py create mode 100644 tests/test_cleanup.py create mode 100644 tests/test_config.py create mode 100644 tests/test_configops.py create mode 100644 tests/test_events.py create mode 100644 tests/test_exceptions.py create mode 100644 tests/test_executor.py create mode 100644 tests/test_exporter.py create mode 100644 tests/test_generator.py create mode 100644 tests/test_inventory.py create mode 100644 tests/test_jump.py create mode 100644 tests/test_logutil.py create mode 100644 tests/test_models.py create mode 100644 tests/test_parser.py create mode 100644 tests/test_permissions.py create mode 100644 tests/test_resolver.py create mode 100644 tests/test_validator.py diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..4aad964 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from sshctl.core.config import SshctlConfig +from sshctl.core.models import HostBlock + + +@pytest.fixture +def ssh_dir(tmp_path: Path) -> Path: + d = tmp_path / ".ssh" + d.mkdir() + return d + + +@pytest.fixture +def sample_config(ssh_dir: Path) -> Path: + content = """Host * + User root + Port 22 + StrictHostKeyChecking accept-new + +Host server1 + HostName 10.0.0.1 + User admin + Port 2222 + +Host server2 + HostName 10.0.0.2 + ProxyJump server1 + +Host *.example.com + HostName example.com + User deploy +""" + path = ssh_dir / "config" + path.write_text(content) + return path + + +@pytest.fixture +def config_with_includes(ssh_dir: Path) -> Path: + conf_d = ssh_dir / "conf.d" + conf_d.mkdir() + + jumps = conf_d / "98-jumpservers.conf" + jumps.write_text("""Host gw + HostName 10.0.0.254 + User jumpuser +""") + + creds = conf_d / "99-cred.conf" + creds.write_text("""Host * + IdentityFile ~/.ssh/keys.d/fallback +""") + + main = ssh_dir / "config" + main.write_text("""Host * + User root + Port 22 + +Include conf.d/98-jumpservers.conf +Include conf.d/99-cred.conf + +Host app1 + HostName 10.0.0.10 + ProxyJump gw +""") + return main + + +@pytest.fixture +def sample_host() -> HostBlock: + host = HostBlock(alias="test-box") + host.set("HostName", "10.0.0.1") + host.set("User", "admin") + host.set("Port", "2222") + host.set("ProxyJump", "gw") + return host + + +@pytest.fixture +def cfg() -> SshctlConfig: + return SshctlConfig() diff --git a/tests/test_checker.py b/tests/test_checker.py new file mode 100644 index 0000000..bc98abf --- /dev/null +++ b/tests/test_checker.py @@ -0,0 +1,119 @@ +from __future__ import annotations + +from sshctl.config.resolver import ResolvedHost +from sshctl.operations.checker import HostCheck, check_host + + +def test_host_check_dataclass() -> None: + h = HostCheck( + alias="test", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="", + ) + assert h.alias == "test" + assert h.hostname == "10.0.0.1" + assert h.user == "admin" + assert h.port == 22 + + +def test_host_check_passed_both_ok() -> None: + h = HostCheck( + alias="test", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="", + ) + assert not h.passed + h.dns_ok = True + assert not h.passed + h.tcp_ok = True + assert h.passed + + +def test_host_check_summary_all_ok() -> None: + h = HostCheck( + alias="test", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="", + ) + h.dns_ok = True + h.tcp_ok = True + h.ssh_ok = True + assert h.summary == "DNS/TCP/SSH" + + +def test_host_check_summary_dns_fail() -> None: + h = HostCheck( + alias="test", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="", + ) + h.dns_error = "timeout" + assert h.summary == "DNS!/TCP!/---" + + +def test_host_check_summary_ssh_untested() -> None: + h = HostCheck( + alias="test", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="", + ) + h.dns_ok = True + h.tcp_ok = True + assert h.summary == "DNS/TCP/---" + + +def test_host_check_summary_ssh_error() -> None: + h = HostCheck( + alias="test", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="", + ) + h.dns_ok = True + h.tcp_ok = True + h.ssh_error = "refused" + assert h.summary == "DNS/TCP/SSH!" + + +async def test_check_host_dns_failure() -> None: + host = ResolvedHost( + alias="nonexistent", + hostname="192.0.2.999", + user="root", + port=22, + proxy_jump="", + identity_file="", + options={}, + ) + result = await check_host(host, timeout=1) + assert isinstance(result.rtt_ms, float) + assert isinstance(result.dns_ok, bool) + assert isinstance(result.tcp_ok, bool) + + +async def test_check_host_constructs_correctly() -> None: + host = ResolvedHost( + alias="test-box", + hostname="192.0.2.1", + user="deploy", + port=2222, + proxy_jump="gw", + identity_file="", + options={}, + ) + result = await check_host(host, timeout=2) + assert result.alias == "test-box" + assert result.hostname == "192.0.2.1" + assert result.port == 2222 + assert result.rtt_ms >= 0 diff --git a/tests/test_cleanup.py b/tests/test_cleanup.py new file mode 100644 index 0000000..f6dee35 --- /dev/null +++ b/tests/test_cleanup.py @@ -0,0 +1,91 @@ +from __future__ import annotations + +from pathlib import Path +from unittest.mock import patch + +from sshctl.operations.cleanup import _is_socket_in_use, cleanup_control_sockets + + +def test_cleanup_no_control_dir(ssh_dir: Path) -> None: + result = cleanup_control_sockets(ssh_dir) + assert result == {"total": 0, "removed": 0, "active": 0} + + +def test_cleanup_empty(ssh_dir: Path) -> None: + control = ssh_dir / "control.d" + control.mkdir() + result = cleanup_control_sockets(ssh_dir) + assert result == {"total": 0, "removed": 0, "active": 0} + + +def test_cleanup_removes_stale(ssh_dir: Path) -> None: + control = ssh_dir / "control.d" + control.mkdir() + (control / "abc.sock").write_text("") + (control / "def.sock").write_text("") + + with patch("sshctl.operations.cleanup._is_socket_in_use", return_value=False): + result = cleanup_control_sockets(ssh_dir) + + assert result["total"] == 2 + assert result["removed"] == 2 + assert result["active"] == 0 + assert not (control / "abc.sock").exists() + assert not (control / "def.sock").exists() + + +def test_cleanup_dry_run_keeps_sockets(ssh_dir: Path) -> None: + control = ssh_dir / "control.d" + control.mkdir() + sock = control / "abc.sock" + sock.write_text("") + + with patch("sshctl.operations.cleanup._is_socket_in_use", return_value=False): + result = cleanup_control_sockets(ssh_dir, dry_run=True) + + assert result["total"] == 1 + assert result["removed"] == 1 + assert sock.exists() + + +def test_cleanup_active_not_removed(ssh_dir: Path) -> None: + control = ssh_dir / "control.d" + control.mkdir() + (control / "abc.sock").write_text("") + (control / "def.sock").write_text("") + + def fake_in_use(sock: Path) -> bool: + return sock.name == "abc.sock" + + with patch("sshctl.operations.cleanup._is_socket_in_use", side_effect=fake_in_use): + result = cleanup_control_sockets(ssh_dir) + + assert result["total"] == 2 + assert result["active"] == 1 + assert result["removed"] == 1 + assert (control / "abc.sock").exists() + assert not (control / "def.sock").exists() + + +def test_is_socket_in_use_with_ssh() -> None: + with patch( + "sshctl.operations.cleanup.subprocess.run", + ) as mock_run: + mock_run.return_value.stdout = " 123 ssh user cwd DIR /tmp/socket" + assert _is_socket_in_use(Path("/tmp/x.sock")) is True + + +def test_is_socket_in_use_no_ssh() -> None: + with patch( + "sshctl.operations.cleanup.subprocess.run", + ) as mock_run: + mock_run.return_value.stdout = " 123 vim user cwd DIR /tmp" + assert _is_socket_in_use(Path("/tmp/x.sock")) is False + + +def test_is_socket_in_use_error() -> None: + with patch( + "sshctl.operations.cleanup.subprocess.run", + side_effect=FileNotFoundError, + ): + assert _is_socket_in_use(Path("/tmp/x.sock")) is False diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..8f53b74 --- /dev/null +++ b/tests/test_config.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.core.config import SshctlConfig, load_config + + +def test_default_config() -> None: + config = SshctlConfig() + assert config.ssh_dir == Path("~/.ssh").expanduser() + assert config.config_file == "config" + assert config.verbose is False + assert config.timeout == 10 + + +def test_config_env_prefix() -> None: + assert SshctlConfig.model_config["env_prefix"] == "SSHCTL_" + + +def test_ssh_config_path() -> None: + config = SshctlConfig(ssh_dir=Path("/tmp/.ssh")) + assert config.ssh_config_path() == Path("/tmp/.ssh/config") + + +def test_conf_d_path() -> None: + config = SshctlConfig(ssh_dir=Path("/tmp/.ssh")) + assert config.conf_d_path() == Path("/tmp/.ssh/conf.d") + + +def test_load_config_nonexistent() -> None: + config = load_config(Path("/nonexistent/sshctl.yaml")) + assert isinstance(config, SshctlConfig) diff --git a/tests/test_configops.py b/tests/test_configops.py new file mode 100644 index 0000000..e87d975 --- /dev/null +++ b/tests/test_configops.py @@ -0,0 +1,293 @@ +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", ""] diff --git a/tests/test_events.py b/tests/test_events.py new file mode 100644 index 0000000..9eda3f4 --- /dev/null +++ b/tests/test_events.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.operations.events import Event, flush, log_event, recent + + +def test_event_to_dict() -> None: + ev = Event( + command="host list", + status="ok", + args={"--verbose": "true"}, + result="", + duration_ms=12.5, + ) + d = ev.to_dict() + assert d["cmd"] == "host list" + assert d["status"] == "ok" + assert d["args"] == {"--verbose": "true"} + assert d["duration_ms"] == 12.5 + assert isinstance(d["ts"], float) + + +def test_log_event_and_flush(ssh_dir: Path) -> None: + log_event(Event(command="check", status="ok", duration_ms=1.5)) + flush(ssh_dir) + rows = recent(ssh_dir) + assert len(rows) == 1 + assert rows[0]["cmd"] == "check" + assert rows[0]["status"] == "ok" + + +def test_flush_empty_no_file(ssh_dir: Path) -> None: + flush(ssh_dir) + assert recent(ssh_dir) == [] + + +def test_recent_empty_when_no_file(ssh_dir: Path) -> None: + assert recent(ssh_dir) == [] + + +def test_recent_limits(ssh_dir: Path) -> None: + for i in range(10): + log_event(Event(command=f"cmd{i}", status="ok")) + flush(ssh_dir) + rows = recent(ssh_dir, limit=3) + assert len(rows) == 3 + assert rows[0]["cmd"] == "cmd7" + + +def test_recent_skips_bad_json(ssh_dir: Path) -> None: + report = ssh_dir / "report" + report.mkdir(exist_ok=True) + (report / "events.ndjson").write_text("not-json\n") + assert recent(ssh_dir) == [] + + +def test_flush_writes_valid_ndjson(ssh_dir: Path) -> None: + log_event(Event(command="tree", status="ok")) + log_event(Event(command="diff", status="error")) + flush(ssh_dir) + path = ssh_dir / "report" / "events.ndjson" + lines = path.read_text().strip().splitlines() + assert len(lines) == 2 diff --git a/tests/test_exceptions.py b/tests/test_exceptions.py new file mode 100644 index 0000000..df8db2b --- /dev/null +++ b/tests/test_exceptions.py @@ -0,0 +1,25 @@ +from __future__ import annotations + +from sshctl.core.exceptions import ( + ConfigError, + ParseError, + ResolverError, + SshctlError, +) + + +def test_exception_hierarchy() -> None: + assert issubclass(ConfigError, SshctlError) + assert issubclass(ParseError, SshctlError) + assert issubclass(ResolverError, SshctlError) + + +def test_sshctl_error_is_base() -> None: + assert issubclass(SshctlError, Exception) + + +def test_exception_raised() -> None: + try: + raise ConfigError("test error") + except SshctlError as e: + assert str(e) == "test error" diff --git a/tests/test_executor.py b/tests/test_executor.py new file mode 100644 index 0000000..c520f91 --- /dev/null +++ b/tests/test_executor.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from sshctl.operations.executor import ExecResult + + +def test_exec_result_dataclass() -> None: + r = ExecResult(alias="server1", returncode=0, stdout="ok", stderr="") + assert r.alias == "server1" + assert r.returncode == 0 + assert r.stdout == "ok" + assert r.error == "" + + +def test_exec_result_error() -> None: + r = ExecResult(alias="server1", returncode=None, error="Timeout") + assert r.returncode is None + assert r.error == "Timeout" + + +def test_exec_result_defaults() -> None: + r = ExecResult(alias="test", returncode=0) + assert r.stdout == "" + assert r.stderr == "" + assert r.error == "" diff --git a/tests/test_exporter.py b/tests/test_exporter.py new file mode 100644 index 0000000..13d6ad5 --- /dev/null +++ b/tests/test_exporter.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.config.resolver import ResolvedHost +from sshctl.export.exporter import export_csv, export_csv_to_file + + +def test_export_csv_with_hosts() -> None: + hosts = [ + ResolvedHost( + alias="server1", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="gw", + identity_file="~/.ssh/id_rsa", + options={}, + ), + ResolvedHost( + alias="server2", + hostname="10.0.0.2", + user="root", + port=2222, + proxy_jump="", + identity_file="", + options={}, + ), + ] + + csv_content = export_csv(hosts) + lines = csv_content.strip().splitlines() + assert len(lines) == 3 + assert lines[0] == "Host,HostName,User,Port,ProxyJump,IdentityFile" + assert "server1" in lines[1] + assert "server2" in lines[2] + assert "10.0.0.1" in lines[1] + assert "gw" in lines[1] + + +def test_export_csv_to_file(tmp_path: Path) -> None: + hosts = [ + ResolvedHost( + alias="test", + hostname="10.0.0.1", + user="root", + port=22, + proxy_jump="", + identity_file="", + options={}, + ), + ] + + path = export_csv_to_file(tmp_path / "output.csv", hosts) + assert path.exists() + content = path.read_text() + assert "test" in content + assert "10.0.0.1" in content + + +def test_export_csv_requires_args() -> None: + raised = False + try: + export_csv() + except ValueError: + raised = True + assert raised, "Expected ValueError" diff --git a/tests/test_generator.py b/tests/test_generator.py new file mode 100644 index 0000000..a7057dd --- /dev/null +++ b/tests/test_generator.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.config.generator import generate_config_text, generate_host_block, write_host_file +from sshctl.config.parser import parse_config +from sshctl.core.models import HostBlock, SshConfig + + +def test_generate_host_block() -> None: + host = HostBlock(alias="test") + host.set("HostName", "10.0.0.1") + host.set("User", "admin") + + text = generate_host_block(host) + assert "Host test" in text + assert "HostName 10.0.0.1" in text + assert "User admin" in text + + +def test_generate_config_roundtrip() -> None: + original = """Host * + User root + Port 22 + +Host server1 + HostName 10.0.0.1 + User admin + Port 2222 + +Include conf.d/98-jumpservers.conf +""" + parsed = parse_config(original) + generated = generate_config_text(parsed, include_header=False) + + reparsed = parse_config(generated) + assert len(reparsed.blocks) == len(parsed.blocks) + assert reparsed.blocks[0].alias == parsed.blocks[0].alias + assert reparsed.blocks[1].hostname == parsed.blocks[1].hostname + assert reparsed.blocks[1].user == parsed.blocks[1].user + + +def test_generate_config_header() -> None: + parsed = parse_config("Host test\n HostName 10.0.0.1\n") + text = generate_config_text(parsed, include_header=True) + assert text.startswith("# Generated by sshctl") + + +def test_generate_empty() -> None: + config = SshConfig() + text = generate_config_text(config) + assert "# Generated by sshctl" in text + + +def test_write_host_file(tmp_path: Path) -> None: + hosts = [ + HostBlock(alias="web01", options={"HostName": "10.0.0.1"}), + HostBlock(alias="web02", options={"HostName": "10.0.0.2"}), + ] + path = write_host_file(hosts, tmp_path / "test.conf") + assert path.exists() + content = path.read_text() + assert "Host web01" in content + assert "Host web02" in content + + +def test_option_ordering() -> None: + host = HostBlock(alias="test") + host.set("HostName", "10.0.0.1") + host.set("User", "admin") + host.set("Port", "2222") + host.set("ProxyJump", "gw") + + text = generate_host_block(host) + lines = text.splitlines() + hostname_idx = next(i for i, line in enumerate(lines) if "HostName" in line) + port_idx = next(i for i, line in enumerate(lines) if "Port" in line) + proxy_idx = next(i for i, line in enumerate(lines) if "ProxyJump" in line) + + assert hostname_idx < proxy_idx + assert port_idx < proxy_idx + + +def test_roundtrip_preserves_options() -> None: + config = SshConfig() + config.raw_includes.append("conf.d/*.conf") + + block = HostBlock(alias="app") + block.set("HostName", "app.example.com") + block.set("User", "deploy") + block.set("Port", "2222") + block.set("ProxyJump", "gw") + block.set("ForwardAgent", "yes") + config.blocks.append(block) + + text = generate_config_text(config, include_header=False) + reparsed = parse_config(text) + reblock = reparsed.find("app") + assert reblock is not None + assert reblock.hostname == "app.example.com" + assert reblock.user == "deploy" + assert reblock.proxy_jump == "gw" + assert reblock.get("ForwardAgent") == "yes" diff --git a/tests/test_inventory.py b/tests/test_inventory.py new file mode 100644 index 0000000..cedc49c --- /dev/null +++ b/tests/test_inventory.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.export.inventory import _to_group_name, generate_inventory + + +def test_to_group_name_simple() -> None: + assert _to_group_name("internal") == "internal" + assert _to_group_name("cluster1") == "cluster1" + + +def test_to_group_name_with_hyphens() -> None: + assert _to_group_name("8xx-serverfarm") == "z8xx_serverfarm" + + +def test_to_group_name_starts_with_digit() -> None: + assert _to_group_name("8xx") == "z8xx" + + +def test_generate_inventory_no_ssh_dir(tmp_path: Path) -> None: + result = generate_inventory(tmp_path / "nonexistent") + assert result == {} + + +def test_generate_inventory_empty_config(tmp_path: Path) -> None: + ssh_dir = tmp_path / ".ssh" + ssh_dir.mkdir() + (ssh_dir / "config").write_text("Host *\n User root\n") + result = generate_inventory(ssh_dir) + assert result == {} diff --git a/tests/test_jump.py b/tests/test_jump.py new file mode 100644 index 0000000..e110c1e --- /dev/null +++ b/tests/test_jump.py @@ -0,0 +1,271 @@ +from __future__ import annotations + +import json +from pathlib import Path +from unittest.mock import patch + +from sshctl.operations.jump import ( + add_network, + current_ssid, + is_internal_by_ssid, + is_internal_by_subnet, + known_networks, + list_hosts, + load_state, + off, + on, + save_state, + status, +) + + +def _mk_conf(ssh_dir: Path, name: str, text: str) -> Path: + conf_d = ssh_dir / "conf.d" + conf_d.mkdir(parents=True, exist_ok=True) + p = conf_d / name + p.write_text(text) + return p + + +def test_known_networks_defaults(ssh_dir: Path) -> None: + assert known_networks(ssh_dir) == ["MobinNet*", "Mobin*"] + + +def test_known_networks_custom(ssh_dir: Path) -> None: + (ssh_dir / ".internal_ssids").write_text("Home*\n # comment\nOffice\n") + assert known_networks(ssh_dir) == ["Home*", "Office"] + + +def test_add_network_creates_file(ssh_dir: Path) -> None: + add_network(ssh_dir, "MyWiFi") + lines = (ssh_dir / ".internal_ssids").read_text().splitlines() + assert lines[-1] == "MyWiFi" + + +def test_add_network_does_not_duplicate(ssh_dir: Path) -> None: + add_network(ssh_dir, "TestNet") + add_network(ssh_dir, "TestNet") + assert (ssh_dir / ".internal_ssids").read_text().strip().count("TestNet") == 1 + + +def test_current_ssid_detected() -> None: + with patch("sshctl.operations.jump.subprocess.run") as m: + m.return_value.stdout = "yes:MyWiFi\nno:OldNet\n" + assert current_ssid() == "MyWiFi" + + +def test_current_ssid_none() -> None: + with patch("sshctl.operations.jump.subprocess.run") as m: + m.return_value.stdout = "no:OldNet\n" + assert current_ssid() is None + + +def test_current_ssid_subprocess_error() -> None: + with patch("sshctl.operations.jump.subprocess.run", side_effect=FileNotFoundError): + assert current_ssid() is None + + +def test_is_internal_by_ssid_no_wifi(tmp_path: Path) -> None: + with patch("sshctl.operations.jump.current_ssid", return_value=None): + assert is_internal_by_ssid(tmp_path) is None + + +def test_is_internal_by_ssid_match(tmp_path: Path) -> None: + (tmp_path / ".internal_ssids").write_text("MobinNet*\n") + with patch("sshctl.operations.jump.current_ssid", return_value="MobinNet-5G"): + assert is_internal_by_ssid(tmp_path) is True + + +def test_is_internal_by_ssid_no_match(tmp_path: Path) -> None: + (tmp_path / ".internal_ssids").write_text("MobinNet*\n") + with patch("sshctl.operations.jump.current_ssid", return_value="CoffeeShop"): + assert is_internal_by_ssid(tmp_path) is False + + +def test_is_internal_by_ssid_default_patterns(tmp_path: Path) -> None: + with patch("sshctl.operations.jump.current_ssid", return_value="Mobin-Office"): + assert is_internal_by_ssid(tmp_path) is True + + +def test_list_hosts(ssh_dir: Path) -> None: + _mk_conf( + ssh_dir, + "10-servers.conf", + "Host gw\n HostName 10.0.0.254\n\nHost web01\n HostName 10.0.0.1\n ProxyJump gw\n", + ) + hosts = list_hosts(ssh_dir) + assert len(hosts) == 1 + assert hosts[0]["alias"] == "web01" + assert hosts[0]["jump"] == "gw" + assert hosts[0]["file"] == "conf.d/10-servers.conf" + + +def test_list_hosts_empty(ssh_dir: Path) -> None: + assert list_hosts(ssh_dir) == [] + + +def test_list_hosts_trailing_host_with_proxy(ssh_dir: Path) -> None: + _mk_conf( + ssh_dir, + "10-servers.conf", + "Host web01\n HostName 10.0.0.1\n ProxyJump gw\n", + ) + hosts = list_hosts(ssh_dir) + assert len(hosts) == 1 + assert hosts[0]["alias"] == "web01" + + +def test_save_state(ssh_dir: Path) -> None: + _mk_conf( + ssh_dir, + "10-servers.conf", + "Host web01\n HostName 10.0.0.1\n ProxyJump gw\n", + ) + records = save_state(ssh_dir) + assert len(records) == 1 + assert records[0]["file"] == "conf.d/10-servers.conf" + assert records[0]["text"] == " ProxyJump gw" + assert (ssh_dir / ".proxyjump_save").exists() + + +def test_save_state_no_proxy(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n HostName 10.0.0.1\n") + assert save_state(ssh_dir) == [] + + +def test_load_state_missing(ssh_dir: Path) -> None: + assert load_state(ssh_dir) is None + + +def test_load_state_invalid_json(ssh_dir: Path) -> None: + (ssh_dir / ".proxyjump_save").write_text("{not json") + assert load_state(ssh_dir) is None + + +def test_load_state_valid(ssh_dir: Path) -> None: + (ssh_dir / ".proxyjump_save").write_text( + json.dumps([{"file": "conf.d/x.conf", "line": 3, "text": " ProxyJump gw"}]) + ) + records = load_state(ssh_dir) + assert records is not None + assert records[0]["file"] == "conf.d/x.conf" + + +def test_off_noop(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n HostName 10.0.0.1\n") + result = off(ssh_dir) + assert result["status"] == "noop" + + +def test_off_removes_proxy_jump(ssh_dir: Path) -> None: + f = _mk_conf( + ssh_dir, + "10-servers.conf", + "Host web01\n HostName 10.0.0.1\n ProxyJump gw\n", + ) + result = off(ssh_dir) + assert result["status"] == "ok" + assert result["count"] == 1 + assert "ProxyJump" not in f.read_text() + assert (ssh_dir / ".proxyjump_save").exists() + + +def test_off_multiple_files(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + _mk_conf(ssh_dir, "20-servers.conf", "Host web02\n ProxyJump gw\n") + result = off(ssh_dir) + assert result["status"] == "ok" + assert result["count"] == 2 + assert len(result["files"]) == 2 + + +def test_on_noop(ssh_dir: Path) -> None: + result = on(ssh_dir) + assert result["status"] == "noop" + + +def test_on_restores(ssh_dir: Path) -> None: + _mk_conf( + ssh_dir, + "10-servers.conf", + "Host web01\n HostName 10.0.0.1\n ProxyJump gw\n", + ) + off(ssh_dir) + f = ssh_dir / "conf.d" / "10-servers.conf" + + result = on(ssh_dir) + assert result["status"] == "ok" + assert "ProxyJump gw" in f.read_text() + assert not (ssh_dir / ".proxyjump_save").exists() + + +def test_on_restores_at_original_line(ssh_dir: Path) -> None: + conf = "Host web01\n HostName 10.0.0.1\n ProxyJump gw\n\nHost web02\n HostName 10.0.0.2\n" + _mk_conf(ssh_dir, "10-servers.conf", conf) + off(ssh_dir) + on(ssh_dir) + text = (ssh_dir / "conf.d" / "10-servers.conf").read_text() + assert text.index("ProxyJump") < text.index("web02") + + +def test_on_skips_file_with_existing_proxy(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + off(ssh_dir) + (ssh_dir / "conf.d" / "10-servers.conf").write_text("Host web01\n ProxyJump newgw\n") + result = on(ssh_dir) + assert result["files"] == [] + assert "newgw" in (ssh_dir / "conf.d" / "10-servers.conf").read_text() + + +def test_on_missing_file_skipped(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + off(ssh_dir) + (ssh_dir / "conf.d" / "10-servers.conf").unlink() + result = on(ssh_dir) + assert result["status"] == "ok" + + +def test_on_line_appended_when_out_of_range(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + off(ssh_dir) + (ssh_dir / "conf.d" / "10-servers.conf").write_text("Host web01\n") + result = on(ssh_dir) + assert result["status"] == "ok" + assert "ProxyJump gw" in (ssh_dir / "conf.d" / "10-servers.conf").read_text() + + +def test_status(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + info = status(ssh_dir) + assert info["count"] == 1 + assert info["saved"] is False + + +def test_status_saved(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + off(ssh_dir) + info = status(ssh_dir) + assert info["saved"] is True + + +def test_is_internal_by_subnet_no_gateways(ssh_dir: Path) -> None: + reachable, probes = is_internal_by_subnet(ssh_dir) + assert reachable is False + assert probes == [] + + +def test_is_internal_by_subnet_reachable(ssh_dir: Path) -> None: + _mk_conf(ssh_dir, "10-servers.conf", "Host web01\n ProxyJump gw\n") + probes = [{"gateway": "gw", "ip": "10.0.0.254", "port": 22, "reachable": True, "error": ""}] + with patch("sshctl.operations.jump._subnet_reachable", return_value=probes): + reachable, result = is_internal_by_subnet(ssh_dir) + assert reachable is True + assert result == probes + + +def test_jump_gateways_collects_unique(ssh_dir: Path) -> None: + from sshctl.operations.jump import _jump_gateways + + _mk_conf(ssh_dir, "10-servers.conf", "Host a\n ProxyJump gw\n") + _mk_conf(ssh_dir, "20-servers.conf", "Host b\n ProxyJump gw:2222\n ProxyJump other\n") + assert _jump_gateways(ssh_dir) == {"gw", "other"} diff --git a/tests/test_logutil.py b/tests/test_logutil.py new file mode 100644 index 0000000..2281afe --- /dev/null +++ b/tests/test_logutil.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +import structlog + +from sshctl.core.logutil import get_logger, setup_logging + + +def test_setup_logging_plain() -> None: + setup_logging(verbose=False) + assert structlog.is_configured() is True + logger = get_logger("test") + assert logger is not None + + +def test_setup_logging_verbose() -> None: + setup_logging(verbose=True) + assert structlog.is_configured() is True + + +def test_get_logger_cached() -> None: + setup_logging(verbose=True) + first = get_logger("cache-test") + second = get_logger("cache-test") + assert first is second diff --git a/tests/test_models.py b/tests/test_models.py new file mode 100644 index 0000000..bb63f52 --- /dev/null +++ b/tests/test_models.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.core.models import HostBlock, SshConfig, ZonePath + + +def test_host_block_defaults() -> None: + host = HostBlock(alias="test") + assert host.alias == "test" + assert host.options == {} + assert host.hostname is None + assert host.user is None + assert host.port is None + assert host.proxy_jump is None + + +def test_host_block_set_get() -> None: + host = HostBlock(alias="test") + host.set("HostName", "10.0.0.1") + host.set("User", "admin") + assert host.hostname == "10.0.0.1" + assert host.user == "admin" + + +def test_host_block_port() -> None: + host = HostBlock(alias="test") + host.set("Port", "2222") + assert host.port == 2222 + + +def test_host_block_port_none() -> None: + host = HostBlock(alias="test") + assert host.port is None + + +def test_host_block_default_get() -> None: + host = HostBlock(alias="test") + assert host.get("Nonexistent") is None + assert host.get("Nonexistent", "default") == "default" + + +def test_ssh_config_find() -> None: + config = SshConfig() + config.blocks.append(HostBlock(alias="server1")) + config.blocks.append(HostBlock(alias="server2")) + + found = config.find("server1") + assert found is not None + assert found.alias == "server1" + assert config.find("nonexistent") is None + + +def test_ssh_config_aliases() -> None: + config = SshConfig() + config.blocks.append(HostBlock(alias="*")) + config.blocks.append(HostBlock(alias="server1")) + config.blocks.append(HostBlock(alias="server2")) + + assert config.aliases == ["server1", "server2"] + + +def test_ssh_config_aliases_exclude_glob_patterns() -> None: + config = SshConfig() + config.blocks.append(HostBlock(alias="web01")) + config.blocks.append(HostBlock(alias="192.168.*")) + config.blocks.append(HostBlock(alias="10.*")) + config.blocks.append(HostBlock(alias="203.0.113.* 198.51.100.*")) + config.blocks.append(HostBlock(alias="web[0-9]")) + config.blocks.append(HostBlock(alias="db?")) + config.blocks.append(HostBlock(alias="*")) + + assert config.aliases == ["web01"] + + +def test_ssh_config_find_all() -> None: + config = SshConfig() + config.blocks.append(HostBlock(alias="web01")) + config.blocks.append(HostBlock(alias="web02")) + config.blocks.append(HostBlock(alias="db01")) + + result = config.find_all("web*") + assert len(result) == 2 + assert result[0].alias == "web01" + assert result[1].alias == "web02" + + +def test_zone_path_internal() -> None: + zp = ZonePath( + universe="internal", + sector="AZ", + cluster="cluster1", + zone="8xx-serverfarm", + vlan="82-k8s", + ) + expected = Path("conf.d/internal/AZ/cluster1/8xx-serverfarm/82-k8s") + assert zp.to_rel_path() == expected + + +def test_zone_path_external() -> None: + zp = ZonePath(universe="external") + assert zp.to_rel_path() == Path("conf.d/external") + + +def test_zone_path_from_rel_path() -> None: + zp = ZonePath.from_rel_path(Path("conf.d/internal/AZ/cluster1")) + assert zp.universe == "internal" + assert zp.sector == "AZ" + assert zp.cluster == "cluster1" + assert zp.zone is None + + +def test_zone_path_invalid() -> None: + raised = False + try: + ZonePath.from_rel_path(Path("other/file")) + except ValueError: + raised = True + assert raised, "Expected ValueError" + + +def test_zone_path_unknown_universe() -> None: + raised = False + try: + ZonePath.from_rel_path(Path("conf.d/unknown")) + except ValueError: + raised = True + assert raised, "Expected ValueError" diff --git a/tests/test_parser.py b/tests/test_parser.py new file mode 100644 index 0000000..6fb161f --- /dev/null +++ b/tests/test_parser.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.config.parser import parse_config, parse_file +from sshctl.core.exceptions import ParseError + + +def test_parse_empty_config() -> None: + config = parse_config("") + assert config.blocks == [] + assert config.raw_includes == [] + + +def test_parse_single_host() -> None: + text = """Host myserver + HostName 10.0.0.1 + User admin + Port 2222 +""" + config = parse_config(text) + assert len(config.blocks) == 1 + block = config.blocks[0] + assert block.alias == "myserver" + assert block.hostname == "10.0.0.1" + assert block.user == "admin" + assert block.port == 2222 + + +def test_parse_host_star() -> None: + text = """Host * + User root + Port 22 +""" + config = parse_config(text) + assert len(config.blocks) == 1 + assert config.blocks[0].alias == "*" + assert config.blocks[0].user == "root" + + +def test_parse_multiple_hosts() -> None: + text = """Host server1 + HostName 10.0.0.1 + +Host server2 + HostName 10.0.0.2 + ProxyJump server1 +""" + config = parse_config(text) + assert len(config.blocks) == 2 + assert config.blocks[0].alias == "server1" + assert config.blocks[1].alias == "server2" + assert config.blocks[1].proxy_jump == "server1" + + +def test_parse_includes() -> None: + text = """Host * + User root + +Include conf.d/internal/*.conf +Include conf.d/98-jumpservers.conf +""" + config = parse_config(text) + assert len(config.raw_includes) == 2 + assert "conf.d/internal/*.conf" in config.raw_includes + assert "conf.d/98-jumpservers.conf" in config.raw_includes + + +def test_parse_comments_and_blanks() -> None: + text = """# This is a comment + +Host test + HostName example.com + + # indent comment + User tester +""" + config = parse_config(text) + assert len(config.blocks) == 1 + assert config.blocks[0].alias == "test" + + +def test_parse_file_not_found() -> None: + raised = False + try: + parse_file("/nonexistent/ssh_config") + except ParseError: + raised = True + assert raised, "Expected ParseError" + + +def test_parse_file(sample_config: Path) -> None: + config = parse_file(sample_config) + assert len(config.blocks) >= 3 + assert config.find("server1") is not None + assert config.find("server2") is not None + + +def test_parse_quoted_values() -> None: + text = """Host test + IdentityFile "~/.ssh/keys.d/id_ed25519" + ProxyJump 'gw.example.com' +""" + config = parse_config(text) + block = config.blocks[0] + assert block.get("IdentityFile") == "~/.ssh/keys.d/id_ed25519" + assert block.get("ProxyJump") == "gw.example.com" + + +def test_find_all_wildcard() -> None: + text = """Host web01 + HostName 10.0.0.1 + +Host web02 + HostName 10.0.0.2 + +Host db01 + HostName 10.0.0.3 +""" + config = parse_config(text) + webs = config.find_all("web*") + assert len(webs) == 2 + + +def test_aliases_property() -> None: + text = """Host * + User root + +Host server1 + HostName 10.0.0.1 + +Host server2 + HostName 10.0.0.2 +""" + config = parse_config(text) + assert config.aliases == ["server1", "server2"] diff --git a/tests/test_permissions.py b/tests/test_permissions.py new file mode 100644 index 0000000..1585071 --- /dev/null +++ b/tests/test_permissions.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.export.permissions import check_permissions, fix_ssh_permissions + + +def test_check_permissions_ok(tmp_path: Path) -> None: + ssh_dir = tmp_path / ".ssh" + ssh_dir.mkdir(mode=0o700) + (ssh_dir / "config").write_text("Host test\n", encoding="utf-8") + (ssh_dir / "config").chmod(0o600) + (ssh_dir / "keys.p").mkdir(mode=0o700) + (ssh_dir / "keys.p" / "key.pub").write_text("ssh-ed25519 AAA...") + (ssh_dir / "keys.p" / "key.pub").chmod(0o644) + + issues = check_permissions(ssh_dir) + total = sum(len(v) for v in issues.values()) + assert total == 0 + + +def test_fix_permissions(tmp_path: Path) -> None: + ssh_dir = tmp_path / ".ssh" + ssh_dir.mkdir(mode=0o755) + (ssh_dir / "config").write_text("Host test\n", encoding="utf-8") + (ssh_dir / "config").chmod(0o644) + (ssh_dir / "id_rsa").write_text("private key data") + (ssh_dir / "id_rsa").chmod(0o644) + + counts = fix_ssh_permissions(ssh_dir) + assert counts["dirs"] >= 1 + assert counts["private"] >= 1 + + assert (ssh_dir / "config").stat().st_mode & 0o777 == 0o600 + assert (ssh_dir / "id_rsa").stat().st_mode & 0o777 == 0o600 diff --git a/tests/test_resolver.py b/tests/test_resolver.py new file mode 100644 index 0000000..7151a96 --- /dev/null +++ b/tests/test_resolver.py @@ -0,0 +1,101 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from sshctl.config.resolver import ( + ResolvedHost, + _parse_ssh_g_output, + resolve_all_hosts, + resolve_host, +) +from sshctl.core.exceptions import ResolverError + + +def test_resolved_host_namedtuple() -> None: + host = ResolvedHost( + alias="server1", + hostname="10.0.0.1", + user="admin", + port=22, + proxy_jump="gw", + identity_file="~/.ssh/id_rsa", + options={"hostname": "10.0.0.1", "user": "admin"}, + ) + assert host.alias == "server1" + assert host.hostname == "10.0.0.1" + assert host.port == 22 + assert host.proxy_jump == "gw" + + +def test_parse_ssh_g_output_basic() -> None: + text = "hostname 10.0.0.1\nuser admin\nport 2222\n" + opts = _parse_ssh_g_output(text) + assert opts["hostname"] == "10.0.0.1" + assert opts["user"] == "admin" + assert opts["port"] == "2222" + + +def test_parse_ssh_g_output_empty() -> None: + assert _parse_ssh_g_output("") == {} + + +def test_parse_ssh_g_output_blank_lines() -> None: + text = "\n\nhostname test\n\n" + opts = _parse_ssh_g_output(text) + assert opts == {"hostname": "test"} + + +def test_parse_ssh_g_output_first_wins() -> None: + text = "hostname first\nhostname second\n" + opts = _parse_ssh_g_output(text) + assert opts["hostname"] == "first" + + +def test_parse_ssh_g_output_quoted_values() -> None: + text = 'identityfile "~/.ssh/key"\n' + opts = _parse_ssh_g_output(text) + assert opts["identityfile"] == "~/.ssh/key" + + +def test_parse_ssh_g_output_single_word_line() -> None: + text = "invalid\nhostname test\n" + opts = _parse_ssh_g_output(text) + assert opts == {"hostname": "test"} + + +def test_resolve_host_ssh_not_found() -> None: + with pytest.raises(ResolverError, match="Failed to resolve"): + resolve_host("test", "/nonexistent/config", ssh_bin="ssh-nonexistent") + + +def test_resolve_all_hosts_empty_config(tmp_path: Path) -> None: + config_path = tmp_path / "config" + config_path.write_text("Host *\n User root\n") + results = resolve_all_hosts(config_path, ssh_bin="ssh-nonexistent") + assert results == [] + + +def test_resolve_host_uses_correct_cmd() -> None: + import subprocess + + def _fake_run(*_args: object, **_kwargs: object) -> object: + class FakeResult: + returncode = 0 + stdout = "hostname 10.0.0.1\nuser testuser\nport 2222\nproxyjump jumpbox\n" + stderr = "" + + return FakeResult() + + original_run = subprocess.run + subprocess.run = _fake_run # type: ignore[assignment] + try: + host = resolve_host("myserver", "/tmp/fake_config") + assert host.alias == "myserver" + assert host.hostname == "10.0.0.1" + assert host.user == "testuser" + assert host.port == 2222 + assert host.proxy_jump == "jumpbox" + finally: + subprocess.run = original_run diff --git a/tests/test_validator.py b/tests/test_validator.py new file mode 100644 index 0000000..137755d --- /dev/null +++ b/tests/test_validator.py @@ -0,0 +1,21 @@ +from __future__ import annotations + +from pathlib import Path + +from sshctl.config.validator import validate_config + + +def test_validate_nonexistent(tmp_path: Path) -> None: + result = validate_config(tmp_path / "nonexistent") + assert not result.passed + assert any("not found" in i.message for i in result.issues) + + +def test_validate_valid(sample_config: Path) -> None: + result = validate_config(sample_config) + assert result.passed + + +def test_validate_with_includes(config_with_includes: Path) -> None: + result = validate_config(config_with_includes) + assert result.passed