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"] == {} 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 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 result["private_key"] == "priv" 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 = ( "peer: PUBKEY1\n endpoint: 203.0.113.1:51820\n allowed ips: 10.0.0.0/24\n" ) result = wireguard._parse_wg_show(output) assert "PUBKEY1" in result assert result["PUBKEY1"]["endpoint"] == "203.0.113.1:51820" def test_empty_output(self): result = wireguard._parse_wg_show("") assert result == {}