Files
sshctl/tests/test_generator.py
T

104 lines
3.1 KiB
Python
Raw Normal View History

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"