import json from copy import deepcopy from unittest.mock import MagicMock, patch import pytest from lib import wireguard @pytest.fixture def temp_config(tmp_path): original = wireguard.CONFIG_PATH wireguard.CONFIG_PATH = tmp_path / "config.json" yield tmp_path wireguard.CONFIG_PATH = original class TestDefaultConfig: def test_returns_skeleton(self): cfg = wireguard.DEFAULT_CONFIG assert cfg["interface"]["name"] == "wg0" assert cfg["interface"]["listen_port"] == 51820 assert cfg["interface"]["private_key"] == "" assert cfg["peers"] == {} def test_classes_have_phase2_fields(self): cfg = wireguard.DEFAULT_CONFIG for ck, cv in cfg["access_classes"].items(): assert "subnet" in cv, f"Class {ck} missing subnet" assert "listen_port" in cv, f"Class {ck} missing listen_port" assert "lan_access" in cv, f"Class {ck} missing lan_access" class TestGetConfig: def test_returns_default_when_no_file(self, temp_config): cfg = wireguard.get_config() assert cfg["interface"]["name"] == "wg0" assert cfg["peers"] == {} def test_loads_existing_config(self, temp_config): wireguard.CONFIG_PATH.write_text( json.dumps( { "interface": { "name": "wg0", "listen_port": 51820, "private_key": "existing-key", "public_key": "existing-pub", "addresses": ["10.137.0.1/24"], "post_up": None, "post_down": None, }, "peers": {}, } ) ) cfg = wireguard.get_config() assert cfg["interface"]["private_key"] == "existing-key" class TestSaveConfig: def test_save_and_reload(self, temp_config): cfg = deepcopy(wireguard.DEFAULT_CONFIG) cfg["interface"]["listen_port"] = 51821 wireguard.save_config(cfg) loaded = wireguard.get_config() assert loaded["interface"]["listen_port"] == 51821 class TestGenerateKeyPair: @patch("lib.wireguard.run_proc") def test_returns_keypair(self, mock_run): mock_run.side_effect = [ MagicMock(returncode=0, stdout="private-key\n"), MagicMock(returncode=0, stdout="public-key\n"), ] private, public = wireguard.generate_keypair() assert private == "private-key" assert public == "public-key" assert mock_run.call_count == 2 assert mock_run.call_args_list[1].kwargs.get("input") == "private-key" class TestClassHelpers: def test_class_interface_name(self): assert wireguard.get_class_interface_name("full") == "wg-full" assert wireguard.get_class_interface_name("internet") == "wg-internet" assert wireguard.get_class_interface_name("custom") == "wg-custom" def test_class_zone_name(self): assert wireguard.get_class_zone_name("full") == "vpn-full" assert wireguard.get_class_zone_name("internet") == "vpn-internet" def test_class_peers(self): cfg = { "peers": { "alice": {"access_class": "full", "public_key": "pk1"}, "bob": {"access_class": "internet", "public_key": "pk2"}, "carol": {"access_class": None, "public_key": "pk3"}, } } full_peers = wireguard._class_peers(cfg, "full") assert len(full_peers) == 1 assert "alice" in full_peers int_peers = wireguard._class_peers(cfg, "internet") assert len(int_peers) == 1 assert "bob" in int_peers class TestGenerateClassConf: def test_returns_none_when_no_peers(self, temp_config): cfg = deepcopy(wireguard.DEFAULT_CONFIG) cfg["access_classes"]["full"]["private_key"] = "test-key" result = wireguard.generate_class_conf(cfg, "full") assert result is None @patch("lib.wireguard.ENV") def test_renders_template_per_class(self, mock_env, temp_config): mock_tmpl = MagicMock() mock_tmpl.render.return_value = "[Interface]\nPrivateKey = x\n" mock_env.get_template.return_value = mock_tmpl cfg = deepcopy(wireguard.DEFAULT_CONFIG) cfg["access_classes"]["full"]["private_key"] = "test-key" cfg["peers"]["alice"] = { "access_class": "full", "public_key": "pub1", "private_key": "priv1", "allowed_ips": ["0.0.0.0/0"], } result = wireguard.generate_class_conf(cfg, "full") assert result is not None assert mock_tmpl.render.call_count == 1 call_kwargs = mock_tmpl.render.call_args.kwargs assert call_kwargs["interface"]["name"] == "wg-full" assert call_kwargs["interface"]["private_key"] == "test-key" assert "alice" in call_kwargs["peers"] def test_raises_when_no_private_key(self, temp_config): cfg = { "interface": { "name": "wg0", "listen_port": 51820, "private_key": "", "addresses": ["10.137.0.1/24"], }, "access_classes": { "full": { "name": "Full LAN Access", "description": "Full access", "subnet": "10.137.0.0/24", "listen_port": 51820, "lan_access": True, }, }, "peers": { "alice": { "access_class": "full", "public_key": "pub1", }, }, } wireguard.save_config(cfg) cfg = wireguard.get_config() with pytest.raises(ValueError, match="no private key"): wireguard.generate_class_conf(cfg, "full") class TestGetPeers: def test_empty_peers(self, temp_config): peers = wireguard.get_peers() assert peers == [] def test_lists_peers_without_private_keys(self, temp_config): cfg = { "interface": { "name": "wg0", "listen_port": 51820, "private_key": "", "public_key": "", "addresses": ["10.137.0.1/24"], "post_up": None, "post_down": None, }, "peers": { "client1": { "public_key": "pub1", "private_key": "priv1", "endpoint": "203.0.113.1:51820", "allowed_ips": ["0.0.0.0/0"], "persistent_keepalive": None, "preshared_key": None, } }, } wireguard.CONFIG_PATH.write_text(json.dumps(cfg)) peers = wireguard.get_peers() assert len(peers) == 1 assert peers[0]["name"] == "client1" assert "private_key" not in peers[0] class TestAddPeer: @patch("lib.wireguard.generate_keypair") def test_adds_new_peer(self, mock_gen, temp_config): mock_gen.return_value = ("priv", "pub") result = wireguard.add_peer("client1", allowed_ips=["10.0.0.0/24"]) assert result["public_key"] == "pub" assert "private_key" not in result assert result["allowed_ips"] == ["10.0.0.0/24"] @patch("lib.wireguard.generate_keypair") def test_updates_existing_peer(self, mock_gen, temp_config): mock_gen.return_value = ("priv", "pub") wireguard.add_peer("client1") wireguard.add_peer("client1", endpoint="203.0.113.1:51820") cfg = wireguard.get_config() assert cfg["peers"]["client1"]["endpoint"] == "203.0.113.1:51820" class TestRemovePeer: @patch("lib.wireguard.generate_keypair") def test_removes_peer(self, mock_gen, temp_config): mock_gen.return_value = ("priv", "pub") wireguard.add_peer("client1") wireguard.remove_peer("client1") cfg = wireguard.get_config() assert "client1" not in cfg["peers"] class TestSetListenPort: def test_set_valid_port(self, temp_config): wireguard.set_listen_port(12345) cfg = wireguard.get_config() assert cfg["interface"]["listen_port"] == 12345 def test_set_invalid_port_raises(self, temp_config): with pytest.raises(ValueError): wireguard.set_listen_port(0) with pytest.raises(ValueError): wireguard.set_listen_port(70000) class TestSetPostHooks: def test_set_post_up(self, temp_config): wireguard.set_post_up("iptables -I FORWARD -i wg0 -j ACCEPT") cfg = wireguard.get_config() assert cfg["interface"]["post_up"] == "iptables -I FORWARD -i wg0 -j ACCEPT" def test_clear_post_up(self, temp_config): wireguard.set_post_up("some-cmd") wireguard.set_post_up(None) cfg = wireguard.get_config() assert cfg["interface"]["post_up"] is None class TestInitialize: @patch("lib.wireguard.generate_keypair") def test_initializes_once(self, mock_gen, temp_config): mock_gen.return_value = ("priv", "pub") cfg = wireguard.initialize() assert cfg["interface"]["private_key"] == "priv" assert cfg["interface"]["public_key"] == "pub" @patch("lib.wireguard.generate_keypair") def test_does_not_overwrite_existing(self, mock_gen, temp_config): existing = { "interface": { "name": "wg0", "listen_port": 51820, "private_key": "original-private", "public_key": "original-pub", "addresses": ["10.137.0.1/24"], "post_up": None, "post_down": None, }, "peers": {}, } wireguard.CONFIG_PATH.write_text(json.dumps(existing)) cfg = wireguard.initialize() assert cfg["interface"]["private_key"] == "original-private" mock_gen.assert_not_called() class TestStatus: @patch("lib.wireguard.run_proc") def test_returns_down_when_interface_down(self, mock_run, temp_config): mock_run.return_value = MagicMock( returncode=1, stdout="", stderr="interface not found" ) result = wireguard.status() assert result["up"] is False @patch("lib.wireguard.run_proc") def test_parses_interface_info(self, mock_run, temp_config): mock_run.return_value = MagicMock( returncode=0, stdout=("interface:\n public key: ABCDEF\n listening port: 51820\n"), ) result = wireguard.status() assert result["up"] is True assert result["interface"]["public_key"] == "ABCDEF" assert result["interface"]["listen_port"] == 51820 @patch("lib.wireguard.run_proc") def test_parses_peer_info(self, mock_run, temp_config): mock_run.return_value = MagicMock( returncode=0, stdout=( "interface:\n" " public key: PUB\n" " listening port: 51820\n" "\n" "peer: PUBKEY1\n" " endpoint: 203.0.113.1:51820\n" " allowed ips: 10.137.0.2/32\n" " latest handshake: 2 minutes ago\n" " transfer: 1.23 GiB received, 4.56 GiB sent\n" " persistent-keepalive: 25\n" ), ) result = wireguard.status() assert len(result["peers"]) == 1 peer = result["peers"][0] assert peer["public_key"] == "PUBKEY1" assert peer["endpoint"] == "203.0.113.1:51820" assert peer["persistent_keepalive"] == 25 class TestGenerateWgShowParser: def test_parses_peer_output(self): output = ( "interface: wg0\n" " public key: IFACE-PUB\n" " listening port: 51820\n" " peer: PUBKEY1\n endpoint: 203.0.113.1:51820\n allowed ips: 10.0.0.0/24\n" ) result = wireguard.parse_wg_show_output(output) assert result["up"] is True assert result["interface"]["public_key"] == "IFACE-PUB" assert result["interface"]["listen_port"] == 51820 assert len(result["peers"]) == 1 assert result["peers"][0]["public_key"] == "PUBKEY1" assert result["peers"][0]["endpoint"] == "203.0.113.1:51820" assert result["peers"][0]["allowed_ips"] == ["10.0.0.0/24"] def test_empty_output(self): result = wireguard.parse_wg_show_output("") assert result["up"] is False assert result["peers"] == [] def test_parses_fwmark(self): output = ( "interface: wg0\n" " public key: IFACE-PUB\n" " listening port: 51820\n" " fwmark: 0x0\n" ) result = wireguard.parse_wg_show_output(output) assert result["up"] is True assert result["interface"]["fwmark"] == "0x0" def test_peer_transfer_and_keepalive(self): output = ( "interface: wg0\n" " public key: IFACE-PUB\n" " listening port: 51820\n" " peer: PUBKEY1\n" " endpoint: 203.0.113.1:51820\n" " allowed ips: 10.0.0.0/24, 10.0.1.0/24\n" " latest handshake: 2 minutes ago\n" " transfer: 1.23 GiB received, 4.56 GiB sent\n" " persistent-keepalive: 25\n" ) result = wireguard.parse_wg_show_output(output) peer = result["peers"][0] assert peer["allowed_ips"] == ["10.0.0.0/24", "10.0.1.0/24"] assert peer["latest_handshake"] == "2 minutes ago" assert peer["transfer_received"] == "1.23 GiB received" assert peer["transfer_sent"] == "4.56 GiB sent" assert peer["persistent_keepalive"] == 25 def test_bad_keepalive_value(self): output = "interface: wg0\n peer: PUBKEY1\n persistent-keepalive: bogus\n" result = wireguard.parse_wg_show_output(output) assert result["peers"][0]["persistent_keepalive"] is None class TestAccessClasses: def test_default_config_has_access_classes(self): cfg = wireguard.DEFAULT_CONFIG assert "access_classes" in cfg assert "full" in cfg["access_classes"] assert "internet" in cfg["access_classes"] def test_ensure_access_classes_empty(self, temp_config): cfg = {"interface": {}, "access_classes": {}, "peers": {}} wireguard._ensure_access_classes(cfg) assert "full" in cfg["access_classes"] assert "internet" in cfg["access_classes"] def test_ensure_access_classes_preserves_existing(self, temp_config): cfg = { "interface": {}, "access_classes": {"custom": {"name": "Custom"}}, "peers": {}, } wireguard._ensure_access_classes(cfg) assert "custom" in cfg["access_classes"] assert "full" not in cfg["access_classes"] def test_ensure_class_defaults_adds_phase2_fields(self, temp_config): cfg = {"interface": {}, "access_classes": {"old": {"name": "Old"}}, "peers": {}} wireguard._ensure_access_classes(cfg) c = cfg["access_classes"]["old"] assert "subnet" in c assert "listen_port" in c assert "lan_access" in c assert "private_key" in c assert "public_key" in c class TestAddPeerWithNewFields: @patch("lib.wireguard.generate_keypair") def test_add_peer_with_description_and_access_class(self, mock_gen, temp_config): mock_gen.return_value = ("priv", "pub") result = wireguard.add_peer( "test-peer", description="Test peer", access_class="full", ) assert result["description"] == "Test peer" assert result["access_class"] == "full" cfg = wireguard.get_config() assert cfg["peers"]["test-peer"]["description"] == "Test peer" assert cfg["peers"]["test-peer"]["access_class"] == "full" @patch("lib.wireguard.generate_keypair") def test_update_peer_preserves_existing_fields(self, mock_gen, temp_config): mock_gen.return_value = ("priv", "pub") wireguard.add_peer("p1", description="original", access_class="full") wireguard.add_peer("p1", endpoint="1.2.3.4:51820") cfg = wireguard.get_config() assert cfg["peers"]["p1"]["endpoint"] == "1.2.3.4:51820" assert cfg["peers"]["p1"]["description"] == "original" assert cfg["peers"]["p1"]["access_class"] == "full" class TestInterfaceHasNewFields: def test_default_has_server_endpoint(self): assert wireguard.DEFAULT_CONFIG["interface"].get("server_endpoint") == "" def test_default_has_description(self): assert wireguard.DEFAULT_CONFIG["interface"].get("description") == "" class TestClassKeyGeneration: @patch("lib.wireguard.run_proc") def test_generates_keypair_for_class(self, mock_run, temp_config): mock_run.side_effect = [ MagicMock(returncode=0, stdout="class-priv\n"), MagicMock(returncode=0, stdout="class-pub\n"), ] # Create config with a fresh class (no prior keys, no auto-merge from file) cfg = { "interface": { "name": "wg0", "private_key": "existing", "public_key": "pub", }, "access_classes": {"fresh": {"name": "Fresh", "description": "New class"}}, "peers": {}, } wireguard.save_config(cfg) priv, pub = wireguard.generate_class_keypair("fresh") assert priv == "class-priv" assert pub == "class-pub" loaded = wireguard.get_config() assert loaded["access_classes"]["fresh"]["private_key"] == "class-priv" assert loaded["access_classes"]["fresh"]["public_key"] == "class-pub" @patch("lib.wireguard.run_proc") def test_returns_existing_keys(self, mock_run, temp_config): cfg = deepcopy(wireguard.DEFAULT_CONFIG) cfg["access_classes"]["full"]["private_key"] = "existing-priv" cfg["access_classes"]["full"]["public_key"] = "existing-pub" wireguard.save_config(cfg) priv, pub = wireguard.generate_class_keypair("full") assert priv == "existing-priv" assert pub == "existing-pub" mock_run.assert_not_called() def test_raises_for_missing_class(self, temp_config): with pytest.raises(ValueError, match="not found"): wireguard.generate_class_keypair("nonexistent") class TestGetPeerStatusMultiInterface: @patch("lib.wireguard.run_proc") def test_aggregates_peers_across_classes(self, mock_run, temp_config): def side_effect(cmd, **kwargs): iface = cmd[-1] if iface == "wg-full": return MagicMock( returncode=0, stdout="interface:\n public key: FULL-PUB\n\npeer: PUB1\n", ) if iface == "wg-internet": return MagicMock( returncode=0, stdout="interface:\n public key: INT-PUB\n\npeer: PUB2\n", ) # Legacy interface if iface == "wg0": return MagicMock(returncode=1, stdout="") return MagicMock(returncode=1, stdout="") cfg = deepcopy(wireguard.DEFAULT_CONFIG) cfg["access_classes"]["full"]["private_key"] = "full-priv" cfg["access_classes"]["internet"]["private_key"] = "int-priv" wireguard.save_config(cfg) mock_run.side_effect = side_effect peers = wireguard.get_peer_status() assert len(peers) == 2 assert peers[0]["public_key"] == "PUB1" assert peers[0]["access_class"] == "full" assert peers[1]["public_key"] == "PUB2" assert peers[1]["access_class"] == "internet" class TestApplyClass: @patch("lib.wireguard.run") @patch("lib.wireguard.run_proc") @patch("lib.wireguard.ENV") def test_apply_class_writes_and_up( self, mock_env, mock_proc, mock_run, temp_config ): mock_tmpl = MagicMock() mock_tmpl.render.return_value = "[Interface]\nPrivateKey = x\n" mock_env.get_template.return_value = mock_tmpl cfg = deepcopy(wireguard.DEFAULT_CONFIG) cfg["access_classes"]["full"]["private_key"] = "test-key" cfg["peers"]["alice"] = { "access_class": "full", "public_key": "pub1", "private_key": "priv1", "allowed_ips": ["0.0.0.0/0"], } wireguard.save_config(cfg) wireguard.apply_class("full") # Should call wg-quick up mock_run.assert_any_call(["wg-quick", "up", "wg-full"], sudo=True)