"""Unit tests for the Windows driver. Every getter sends one PowerShell script and parses the JSON it emits. A fake transport answers each script with a fixture from tests/fixtures/synthetic/ — see the README there for what those fixtures do and do not prove. """ from __future__ import annotations import json from pathlib import Path import pytest from napalm.base.exceptions import ConnectionClosedException from napalm_device_types import role_keys_of from napalm_windows import WindowsDriver from napalm_windows import windows as mod from napalm_windows.transport import PowerShellError FIXTURES = Path(__file__).parent / "fixtures" / "synthetic" def _fixture(name: str) -> str: return (FIXTURES / name).read_text() class FakeTransport: """Answers the driver's scripts from fixtures and records what it was sent.""" def __init__(self, answers: dict[str, str] | None = None, error: str | None = None): self.answers = answers or {} self.error = error self.sent: list[str] = [] self.is_open = True def run(self, script: str) -> str: self.sent.append(script) if self.error: raise PowerShellError(self.error) return self.answers.get(script, "") def close(self) -> None: self.is_open = False def _driver(**answers: str) -> WindowsDriver: d = WindowsDriver("win01", "admin", "secret") d._transport = FakeTransport({getattr(mod, k): v for k, v in answers.items()}) return d # --------------------------------------------------------------------------- # Construction and class attributes # --------------------------------------------------------------------------- class TestInit: def test_defaults_to_winrm_over_https(self): d = WindowsDriver("win01", "admin", "secret") assert d.port == 5986 assert d.ssl is True assert d.cert_validation is True assert d.auth == "negotiate" def test_port_5985_means_plain_http(self): d = WindowsDriver("win01", "admin", "secret", optional_args={"port": 5985}) assert d.ssl is False def test_winrm_ssl_overrides_the_port_guess(self): d = WindowsDriver( "win01", "admin", "secret", optional_args={"port": 8443, "winrm_ssl": False} ) assert d.ssl is False def test_ssl_verify_false_disables_cert_validation(self): d = WindowsDriver("win01", "admin", "secret", optional_args={"ssl_verify": False}) assert d.cert_validation is False def test_winrm_auth_is_passed_through(self): d = WindowsDriver("win01", "admin", "secret", optional_args={"winrm_auth": "ntlm"}) assert d.auth == "ntlm" def test_construction_does_no_io(self): assert WindowsDriver("win01", "admin", "secret")._transport is None class TestClassAttributes: def test_driver_name_matches_the_entry_point(self): # netork/core/nvd/platform.py already keys on the driver name "windows". assert WindowsDriver.DRIVER_NAME == "windows" def test_does_not_ask_for_ssh_credentials(self): assert WindowsDriver.USES_SSH is False def test_default_port_is_winrm_https(self): assert WindowsDriver.default_port == 5986 def test_fills_the_general_purpose_os_role(self): # OSDriver's role key is "linux" but means "general-purpose OS host" # (poll timeout, OS tabs). Renaming it is tracked in netork#300. assert role_keys_of(WindowsDriver) == ["linux"] def test_wsman_ports_are_probed_during_discovery(self): ports = {(p.scheme, p.port) for p in WindowsDriver.PORT_SPECS or []} assert ("http", 5985) in ports assert ("https", 5986) in ports def test_http_sys_and_iis_server_headers_are_fingerprints(self): patterns = {r.pattern for r in WindowsDriver.HTTP_FINGERPRINT} assert {"microsoft-httpapi", "microsoft-iis"} <= patterns def test_openssh_for_windows_banner_is_a_fingerprint(self): patterns = {r.pattern for r in WindowsDriver.SSH_FINGERPRINT} assert "openssh_for_windows" in patterns # --------------------------------------------------------------------------- # Connection lifecycle # --------------------------------------------------------------------------- class TestLifecycle: def test_is_alive_false_before_open(self): assert WindowsDriver("win01", "admin", "secret").is_alive() == {"is_alive": False} def test_is_alive_follows_the_transport(self): d = _driver() assert d.is_alive() == {"is_alive": True} def test_close_drops_the_transport(self): d = _driver() fake = d._transport d.close() assert fake.is_open is False assert d._transport is None def test_getter_before_open_raises_connection_closed(self): with pytest.raises(ConnectionClosedException): WindowsDriver("win01", "admin", "secret").get_facts() # --------------------------------------------------------------------------- # JSON handling # --------------------------------------------------------------------------- class TestRunPs: def test_empty_output_is_none(self): d = _driver() assert d._run_ps("Get-Nothing") is None def test_output_is_parsed_as_json(self): d = WindowsDriver("win01", "admin", "secret") d._transport = FakeTransport({"Get-X": '{"a": 1}'}) assert d._run_ps("Get-X") == {"a": 1} class TestAsList: """ConvertTo-Json unwraps a one-element array into a bare object.""" def test_none_is_empty(self): assert mod._as_list(None) == [] def test_single_object_is_wrapped(self): assert mod._as_list({"a": 1}) == [{"a": 1}] def test_list_is_unchanged(self): assert mod._as_list([1, 2]) == [1, 2] class TestPsQuote: def test_wraps_in_single_quotes(self): assert mod._ps_quote("Spooler") == "'Spooler'" def test_doubles_ascii_single_quote(self): assert mod._ps_quote("a'b") == "'a''b'" @pytest.mark.parametrize("quote", ["\u2018", "\u2019", "\u201a", "\u201b"]) def test_doubles_typographic_quotes_powershell_also_accepts(self, quote): # PowerShell treats these as single-quote delimiters too; leaving one # undoubled would end the string early. assert mod._ps_quote(f"a{quote}b") == f"'a{quote}{quote}b'" class TestMac: def test_windows_dashes_become_colons(self): assert mod._mac("00-15-5d-01-02-03") == "00:15:5D:01:02:03" def test_empty_stays_empty(self): assert mod._mac("") == "" assert mod._mac(None) == "" # --------------------------------------------------------------------------- # Getters # --------------------------------------------------------------------------- class TestGetFacts: def test_domain_member_server(self): facts = _driver(_PS_FACTS=_fixture("facts_server.json")).get_facts() assert facts == { "hostname": "srv-app01", "fqdn": "srv-app01.corp.example", "vendor": "Microsoft Corporation", "model": "Virtual Machine", "serial_number": "0000-0001-2345-6789-0123-4567-89", "os_version": "Microsoft Windows Server 2022 Standard 21H2 (build 20348.2340)", "uptime": 86400, "interface_list": ["Ethernet", "Ethernet 2"], "running_kernel": "10.0.20348.2340", } def test_workgroup_client_has_no_domain_suffix(self): facts = _driver(_PS_FACTS=_fixture("facts_client.json")).get_facts() assert facts["hostname"] == "DESKTOP-4F2K9" assert facts["fqdn"] == "DESKTOP-4F2K9" def test_single_interface_arrives_as_a_bare_string(self): facts = _driver(_PS_FACTS=_fixture("facts_client.json")).get_facts() assert facts["interface_list"] == ["Wi-Fi"] def test_missing_manufacturer_falls_back_to_microsoft(self): data = json.loads(_fixture("facts_client.json")) data["manufacturer"] = None facts = _driver(_PS_FACTS=json.dumps(data)).get_facts() assert facts["vendor"] == "Microsoft" def test_os_version_without_display_version(self): # Server 2016 has no DisplayVersion registry value. data = json.loads(_fixture("facts_server.json")) data["caption"] = "Microsoft Windows Server 2016 Standard" data["display_version"] = None data["version"] = "10.0.14393" data["ubr"] = 7428 facts = _driver(_PS_FACTS=json.dumps(data)).get_facts() assert facts["os_version"] == "Microsoft Windows Server 2016 Standard (build 14393.7428)" assert facts["running_kernel"] == "10.0.14393.7428" class TestGetInterfaces: def test_maps_adapters(self): ifaces = _driver(_PS_INTERFACES=_fixture("interfaces.json")).get_interfaces() assert ifaces["Ethernet"] == { "is_up": True, "is_enabled": True, "description": "Microsoft Hyper-V Network Adapter", "last_flapped": -1.0, "speed": 10000.0, "mtu": 1500, "mac_address": "00:15:5D:01:02:03", } def test_disconnected_is_enabled_but_down(self): ifaces = _driver(_PS_INTERFACES=_fixture("interfaces.json")).get_interfaces() assert ifaces["Ethernet 2"]["is_up"] is False assert ifaces["Ethernet 2"]["is_enabled"] is True assert ifaces["Ethernet 2"]["speed"] == 0.0 def test_disabled_adapter_with_null_fields(self): ifaces = _driver(_PS_INTERFACES=_fixture("interfaces.json")).get_interfaces() assert ifaces["Wi-Fi"]["is_enabled"] is False assert ifaces["Wi-Fi"]["mtu"] == 0 assert ifaces["Wi-Fi"]["speed"] == 0.0 assert ifaces["Wi-Fi"]["mac_address"] == "" class TestGetInterfacesIp: def test_groups_addresses_by_interface_and_family(self): ips = _driver(_PS_INTERFACES_IP=_fixture("interfaces_ip.json")).get_interfaces_ip() assert ips == { "Ethernet": { "ipv4": { "10.0.0.5": {"prefix_length": 24}, "10.0.0.6": {"prefix_length": 24}, }, "ipv6": { "fe80::1c2d:3e4f:5a6b:7c8d": {"prefix_length": 64}, "2001:db8::5": {"prefix_length": 64}, }, } } class TestGetArpTable: def test_keeps_only_real_neighbours(self): arp = _driver(_PS_ARP=_fixture("arp.json")).get_arp_table() assert arp == [ {"interface": "Ethernet", "mac": "00:0D:B9:11:22:33", "ip": "10.0.0.1", "age": 0.0}, {"interface": "Ethernet", "mac": "00:15:5D:AA:BB:CC", "ip": "10.0.0.20", "age": 0.0}, ] class TestGetRouteTo: def test_maps_protocols_and_drops_noise(self): routes = _driver(_PS_ROUTES=_fixture("routes.json")).get_route_to() assert set(routes) == {"0.0.0.0/0", "10.0.0.0/24", "192.168.50.0/24", "2001:db8::/64"} assert routes["0.0.0.0/0"][0]["protocol"] == "static" assert routes["10.0.0.0/24"][0]["protocol"] == "connected" assert routes["192.168.50.0/24"][0]["protocol"] == "dhcp" assert routes["2001:db8::/64"][0]["protocol"] == "connected" def test_entry_shape(self): routes = _driver(_PS_ROUTES=_fixture("routes.json")).get_route_to() assert routes["192.168.50.0/24"] == [ { "protocol": "dhcp", "family": "ipv4", "current_active": True, "last_active": False, "age": -1, "next_hop": "10.0.0.254", "outgoing_interface": "Ethernet", "selected_next_hop": True, "preference": 10, "routing_table": "global", "protocol_attributes": {}, } ] def test_on_link_next_hop_is_empty(self): routes = _driver(_PS_ROUTES=_fixture("routes.json")).get_route_to() assert routes["10.0.0.0/24"][0]["next_hop"] == "" assert routes["2001:db8::/64"][0]["next_hop"] == "" assert routes["2001:db8::/64"][0]["family"] == "ipv6" def test_destination_filter(self): routes = _driver(_PS_ROUTES=_fixture("routes.json")).get_route_to(destination="0.0.0.0/0") assert list(routes) == ["0.0.0.0/0"] def test_protocol_filter(self): routes = _driver(_PS_ROUTES=_fixture("routes.json")).get_route_to(protocol="dhcp") assert list(routes) == ["192.168.50.0/24"] class TestGetServices: def test_maps_state_and_start_mode(self): services = _driver(_PS_SERVICES=_fixture("services.json")).get_services() assert services == [ {"name": "WinRM", "running": True, "enabled": True, "pid": 1234}, {"name": "Spooler", "running": False, "enabled": False, "pid": 0}, {"name": "MSSQL$SQLEXPRESS", "running": True, "enabled": True, "pid": 4321}, {"name": "RemoteRegistry", "running": False, "enabled": False, "pid": 0}, {"name": "wuauserv", "running": False, "enabled": False, "pid": 0}, ] class TestManageService: @pytest.mark.parametrize( ("action", "command"), [ ("start", "Start-Service -Name 'Spooler'"), ("stop", "Stop-Service -Name 'Spooler'"), ("restart", "Restart-Service -Name 'Spooler'"), ("enable", "Set-Service -Name 'Spooler' -StartupType Automatic"), ("disable", "Set-Service -Name 'Spooler' -StartupType Disabled"), ], ) def test_sends_the_matching_cmdlet(self, action, command): d = _driver() result = d.manage_service("Spooler", action) assert result["success"] is True assert command in d._transport.sent[0] assert "-ErrorAction Stop" in d._transport.sent[0] def test_service_name_with_dollar_is_valid(self): d = _driver() assert d.manage_service("MSSQL$SQLEXPRESS", "restart")["success"] is True assert "'MSSQL$SQLEXPRESS'" in d._transport.sent[0] def test_unknown_action_is_rejected(self): with pytest.raises(ValueError, match="action"): _driver().manage_service("Spooler", "reload") @pytest.mark.parametrize("name", ["", "a'; Remove-Item C:\\ -Recurse", "a b", "a`b"]) def test_suspicious_name_is_rejected_before_anything_is_sent(self, name): d = _driver() with pytest.raises(ValueError, match="service name"): d.manage_service(name, "start") assert d._transport.sent == [] def test_powershell_error_is_reported_not_raised(self): d = WindowsDriver("win01", "admin", "secret") d._transport = FakeTransport(error="Cannot find any service with service name 'nope'.") result = d.manage_service("nope", "start") assert result == { "success": False, "output": "Cannot find any service with service name 'nope'.", }