#!/usr/bin/env python3
"""
Tests for zylo-agent's pure logic.

Covers the parts that can be verified without a WireGuard interface: parsing
`wg show dump`, the reconciliation diff, /proc metric parsing, and config
validation. These are where the bugs that matter live — a mis-parsed dump field
silently reports the wrong customer's usage, and a reconciliation diff that
misses a change leaves a revoked peer connected.

Run:  python3 test_agent.py
"""

import importlib.machinery
import importlib.util
import os
import sys
import tempfile
import unittest
import unittest.mock
from pathlib import Path
from unittest.mock import patch

# The agent ships without a .py extension (it is an executable), so load it by
# path rather than by import name.
#
# It must be registered in sys.modules *before* exec_module: @dataclass resolves
# type annotations via sys.modules[cls.__module__], which is None for a module
# that is still being constructed.
_spec = importlib.util.spec_from_loader(
    "zylo_agent",
    importlib.machinery.SourceFileLoader("zylo_agent", str(Path(__file__).parent / "zylo-agent")),
)
agent = importlib.util.module_from_spec(_spec)
sys.modules["zylo_agent"] = agent
_spec.loader.exec_module(agent)


class WireGuardDumpParsingTest(unittest.TestCase):
    """`wg show dump` is tab-separated with no header. Field order is the
    contract; getting it wrong misattributes usage between customers."""

    # interface line: private-key, public-key, listen-port, fwmark
    # peer lines: public-key, psk, endpoint, allowed-ips, handshake, rx, tx, keepalive
    DUMP = (
        "aPrivateKey\taPublicKey\t51820\toff\n"
        "peerKeyOne=\t(none)\t203.0.113.5:1234\t10.10.0.2/32\t1717000000\t1024\t2048\t25\n"
        "peerKeyTwo=\tpresharedAAA=\t(none)\t10.10.0.3/32\t0\t0\t0\toff\n"
    )

    def setUp(self):
        self.wg = agent.WireGuard("wg0")

    def test_parses_peers_and_normalises_none_psk(self):
        with patch.object(agent, "run", return_value=self.DUMP):
            peers = self.wg.current_peers()

        self.assertEqual({"peerKeyOne=", "peerKeyTwo="}, set(peers))
        self.assertEqual("10.10.0.2/32", peers["peerKeyOne="].allowed_ips)

        # "(none)" is a literal wg prints, not a preshared key.
        self.assertIsNone(peers["peerKeyOne="].preshared_key)
        self.assertEqual("presharedAAA=", peers["peerKeyTwo="].preshared_key)

    def test_ignores_the_interface_line(self):
        with patch.object(agent, "run", return_value=self.DUMP):
            self.assertNotIn("aPublicKey", self.wg.current_peers())

    def test_parses_transfer_counters(self):
        with patch.object(agent, "run", return_value=self.DUMP):
            entries = {e["public_key"]: e for e in self.wg.transfer()}

        self.assertEqual(1024, entries["peerKeyOne="]["rx_bytes"])
        self.assertEqual(2048, entries["peerKeyOne="]["tx_bytes"])
        self.assertEqual(1717000000, entries["peerKeyOne="]["last_handshake"])

    def test_survives_a_truncated_dump(self):
        # A malformed line must not take the agent down mid-cycle.
        with patch.object(agent, "run", return_value="iface\tline\t1\toff\nbroken\n"):
            self.assertEqual({}, self.wg.current_peers())
            self.assertEqual([], self.wg.transfer())

    def test_returns_empty_when_the_interface_is_absent(self):
        with patch.object(agent, "run", side_effect=RuntimeError("No such device")):
            self.assertEqual({}, self.wg.current_peers())
            self.assertEqual([], self.wg.transfer())


class ReconciliationTest(unittest.TestCase):
    """The diff that decides what actually changes on the interface."""

    def setUp(self):
        config = agent.Config(panel_url="https://panel.test", token="t")
        self.agent = agent.Agent(
            config=config,
            client=agent.PanelClient(config),
            wg=agent.WireGuard("wg0"),
        )

    def _sync(self, current, desired_peers):
        """Runs sync() against a faked interface, returning the calls made."""
        calls = {"added": [], "removed": []}

        state = {
            "server": {"address": "10.10.0.1/24", "listen_port": 51820, "mtu": 1420},
            "peers": desired_peers,
            "state_hash": "a" * 64,
        }

        with patch.object(self.agent.client, "state", return_value=state), \
             patch.object(self.agent.client, "acknowledge", return_value={}), \
             patch.object(self.agent.wg, "ensure_interface"), \
             patch.object(self.agent.wg, "current_peers", return_value=current), \
             patch.object(self.agent.wg, "add_peer", side_effect=lambda p: calls["added"].append(p.public_key)), \
             patch.object(self.agent.wg, "remove_peer", side_effect=lambda k: calls["removed"].append(k)):
            self.agent.sync()

        return calls

    def test_adds_a_new_peer(self):
        calls = self._sync(
            current={},
            desired_peers=[{"public_key": "new=", "allowed_ips": "10.10.0.2/32", "preshared_key": None}],
        )

        self.assertEqual(["new="], calls["added"])
        self.assertEqual([], calls["removed"])

    def test_removes_a_peer_that_is_no_longer_desired(self):
        # This is revocation actually taking effect.
        calls = self._sync(
            current={"gone=": agent.Peer("gone=", "10.10.0.9/32")},
            desired_peers=[],
        )

        self.assertEqual(["gone="], calls["removed"])
        self.assertEqual([], calls["added"])

    def test_makes_no_changes_when_state_already_matches(self):
        peer = agent.Peer("same=", "10.10.0.2/32")

        calls = self._sync(
            current={"same=": peer},
            desired_peers=[{"public_key": "same=", "allowed_ips": "10.10.0.2/32", "preshared_key": None}],
        )

        self.assertEqual([], calls["added"], "an unchanged peer must not be rewritten")
        self.assertEqual([], calls["removed"])

    def test_updates_a_peer_whose_allowed_ips_changed(self):
        calls = self._sync(
            current={"k=": agent.Peer("k=", "10.10.0.2/32")},
            desired_peers=[{"public_key": "k=", "allowed_ips": "10.10.0.7/32", "preshared_key": None}],
        )

        self.assertEqual(["k="], calls["added"], "wg set updates an existing peer in place")
        self.assertEqual([], calls["removed"])

    def test_updates_a_peer_whose_preshared_key_changed(self):
        calls = self._sync(
            current={"k=": agent.Peer("k=", "10.10.0.2/32", preshared_key=None)},
            desired_peers=[{"public_key": "k=", "allowed_ips": "10.10.0.2/32", "preshared_key": "psk="}],
        )

        self.assertEqual(["k="], calls["added"])

    def test_records_the_applied_hash_after_a_successful_sync(self):
        self._sync(current={}, desired_peers=[])

        # Without this the agent re-syncs every poll forever.
        self.assertEqual("a" * 64, self.agent.applied_hash)

    def test_does_not_touch_the_interface_when_no_address_is_provisioned(self):
        state = {"server": {"address": None}, "peers": [], "state_hash": "b" * 64}

        with patch.object(self.agent.client, "state", return_value=state), \
             patch.object(self.agent.wg, "ensure_interface") as ensure:
            self.agent.sync()

        # Bringing an interface up with no address accepts handshakes and then
        # blackholes every packet — worse than staying down.
        ensure.assert_not_called()
        self.assertIsNone(self.agent.applied_hash)


class SystemStatsTest(unittest.TestCase):
    def test_cpu_returns_none_on_first_sample(self):
        stats = agent.SystemStats()
        proc_stat = "cpu  100 0 100 800 0 0 0 0 0 0\n"

        with patch("builtins.open", unittest.mock.mock_open(read_data=proc_stat)):
            # Utilisation is a rate; there is no interval yet. Reporting 0.0
            # would look like an idle machine rather than a missing sample.
            self.assertIsNone(stats.cpu_percent())

    def test_cpu_computes_utilisation_across_samples(self):
        stats = agent.SystemStats()

        with patch("builtins.open", unittest.mock.mock_open(read_data="cpu  100 0 100 800 0 0 0 0 0 0\n")):
            stats.cpu_percent()

        # 100 more busy ticks, 100 more idle ticks => 50%.
        with patch("builtins.open", unittest.mock.mock_open(read_data="cpu  200 0 100 900 0 0 0 0 0 0\n")):
            self.assertEqual(50.0, stats.cpu_percent())

    def test_memory_uses_available_not_free(self):
        meminfo = "MemTotal: 1000 kB\nMemFree: 100 kB\nMemAvailable: 800 kB\n"

        with patch("builtins.open", unittest.mock.mock_open(read_data=meminfo)):
            # MemFree alone would report 90% used on a healthy box whose cache
            # is reclaimable.
            self.assertEqual(20.0, agent.SystemStats.memory_percent())

    def test_metrics_degrade_to_none_rather_than_crashing(self):
        with patch("builtins.open", side_effect=OSError):
            self.assertIsNone(agent.SystemStats.memory_percent())
            self.assertIsNone(agent.SystemStats.uptime_seconds())

        self.assertEqual((0, 0), agent.SystemStats.interface_bytes("nonexistent0"))


class ConfigTest(unittest.TestCase):
    def _write(self, body: str) -> str:
        handle = tempfile.NamedTemporaryFile("w", suffix=".conf", delete=False)
        handle.write(body)
        handle.close()
        self.addCleanup(os.unlink, handle.name)
        return handle.name

    def test_loads_a_valid_config(self):
        path = self._write(
            "[agent]\npanel_url = https://panel.test/\ntoken = zylo_node_abc\ninterface = wg1\n"
        )
        config = agent.Config.load(path)

        self.assertEqual("https://panel.test", config.panel_url, "trailing slash is stripped")
        self.assertEqual("wg1", config.interface)

    def test_rejects_plaintext_http_by_default(self):
        # http:// would put the agent token on the wire in clear every poll.
        path = self._write("[agent]\npanel_url = http://panel.test\ntoken = t\n")

        with self.assertRaises(SystemExit) as caught:
            agent.Config.load(path)

        self.assertIn("cleartext", str(caught.exception))

    def test_allows_http_when_explicitly_acknowledged(self):
        path = self._write(
            "[agent]\npanel_url = http://panel.test\ntoken = t\nallow_insecure = true\n"
        )
        self.assertEqual("http://panel.test", agent.Config.load(path).panel_url)

    def test_requires_a_token(self):
        path = self._write("[agent]\npanel_url = https://panel.test\n")

        with self.assertRaises(SystemExit):
            agent.Config.load(path)

    def test_rejects_a_missing_file(self):
        with self.assertRaises(SystemExit):
            agent.Config.load("/nonexistent/agent.conf")


class HeartbeatPayloadTest(unittest.TestCase):
    def test_omits_metrics_it_could_not_read(self):
        config = agent.Config(panel_url="https://panel.test", token="t")
        instance = agent.Agent(
            config=config,
            client=agent.PanelClient(config),
            wg=agent.WireGuard("wg0"),
        )

        with patch.object(instance.wg, "current_peers", return_value={}), \
             patch.object(instance.wg, "exists", return_value=False), \
             patch.object(instance.wg, "public_key", return_value=None), \
             patch.object(agent.SystemStats, "memory_percent", return_value=None), \
             patch.object(agent.SystemStats, "disk_percent", return_value=None), \
             patch.object(agent.SystemStats, "uptime_seconds", return_value=None):
            payload = instance.build_heartbeat()

        # Sending nulls would have the panel record them as real samples.
        self.assertNotIn("ram_usage", payload)
        self.assertNotIn("disk_usage", payload)
        self.assertNotIn("public_key", payload)

        # wireguard_running is False here, which is meaningful and must survive
        # the None-stripping filter.
        self.assertIs(False, payload["wireguard_running"])
        self.assertEqual(0, payload["active_peers"])


if __name__ == "__main__":
    unittest.main(verbosity=2)
