#!/usr/bin/env python3
"""
zylo-agent — ZyloVPN node agent.

Polls the ZyloVPN control plane for the peer set this node should carry and
reconciles the local WireGuard interface to match.

Design notes
------------
* **Pull, not push.** The control plane never connects to this node. That means
  the only inbound port a ZyloVPN node needs is the WireGuard UDP port — no
  HTTPS listener, no public certificate, and it works behind NAT.

* **Declarative, not imperative.** The agent asks "what should my peers be?"
  and converges. It does not replay a queue of add/remove commands. A node that
  reboots, is restored from a snapshot, or is hand-edited with `wg set` heals
  itself on the next poll instead of drifting forever.

* **`wg set`, never `wg-quick`.** Peer changes are applied to the live
  interface. Tearing the interface down and back up would drop every
  established session on the node to add one peer.

* **Stdlib only.** No pip install on a machine carrying customer traffic. Ships
  as one file; Ubuntu 24.04's Python 3.12 runs it as-is.

Exit codes: 0 clean shutdown, 1 fatal configuration error, 2 missing dependency.
"""

from __future__ import annotations

import configparser
import json
import logging
import os
import signal
import socket
import subprocess
import sys
import tempfile
import time
import urllib.error
import urllib.request
from dataclasses import dataclass, field
from typing import Any

VERSION = "1.0.0"
DEFAULT_CONFIG = "/etc/zylovpn/agent.conf"
USER_AGENT = f"zylo-agent/{VERSION}"

log = logging.getLogger("zylo-agent")


# ---------------------------------------------------------------- configuration


@dataclass
class Config:
    panel_url: str
    token: str
    interface: str = "wg0"
    poll_interval: int = 5
    usage_interval: int = 60
    timeout: int = 15
    verify_tls: bool = True

    @classmethod
    def load(cls, path: str) -> "Config":
        parser = configparser.ConfigParser()

        if not parser.read(path):
            raise SystemExit(f"Cannot read config file: {path}")

        if not parser.has_section("agent"):
            raise SystemExit(f"Config file {path} has no [agent] section.")

        section = parser["agent"]

        panel = section.get("panel_url", "").strip().rstrip("/")
        token = section.get("token", "").strip()

        if not panel:
            raise SystemExit("panel_url is required in [agent].")
        if not token:
            raise SystemExit("token is required in [agent].")

        # A plaintext panel URL would put the agent token on the wire in clear
        # on every poll. Allowed only when explicitly acknowledged, so it can be
        # used for local testing but never reached by accident in production.
        if panel.startswith("http://") and not section.getboolean("allow_insecure", False):
            raise SystemExit(
                "panel_url uses http://, which would transmit the agent token in "
                "cleartext. Use https://, or set allow_insecure = true if this is "
                "a local test environment."
            )

        return cls(
            panel_url=panel,
            token=token,
            interface=section.get("interface", "wg0").strip(),
            poll_interval=section.getint("poll_interval", 5),
            usage_interval=section.getint("usage_interval", 60),
            timeout=section.getint("timeout", 15),
            verify_tls=section.getboolean("verify_tls", True),
        )


# --------------------------------------------------------------------- helpers


def run(args: list[str], *, check: bool = True, stdin: str | None = None) -> str:
    """Runs a command and returns stdout.

    Arguments are passed as a list, never through a shell, so a hostile value
    arriving from the API cannot become shell metacharacters.
    """
    log.debug("exec: %s", " ".join(args))

    result = subprocess.run(
        args,
        capture_output=True,
        text=True,
        input=stdin,
        check=False,
    )

    if check and result.returncode != 0:
        raise RuntimeError(
            f"{args[0]} failed ({result.returncode}): {result.stderr.strip()}"
        )

    return result.stdout


def require_binaries() -> None:
    for binary in ("wg", "ip"):
        if subprocess.run(["which", binary], capture_output=True).returncode != 0:
            log.error("Required binary '%s' not found. Install wireguard-tools.", binary)
            sys.exit(2)


# ------------------------------------------------------------------ system stats


class SystemStats:
    """Reads host metrics from /proc and the filesystem. No dependencies."""

    def __init__(self) -> None:
        self._last_cpu: tuple[int, int] | None = None

    def cpu_percent(self) -> float | None:
        """CPU utilisation since the previous call.

        Returns None on the first call: utilisation is a rate, and there is no
        interval to measure against yet. Reporting 0.0 would look like an idle
        machine rather than a missing sample.
        """
        try:
            with open("/proc/stat") as handle:
                fields = [int(v) for v in handle.readline().split()[1:]]
        except OSError:
            return None

        idle = fields[3] + (fields[4] if len(fields) > 4 else 0)
        total = sum(fields)

        previous, self._last_cpu = self._last_cpu, (idle, total)

        if previous is None:
            return None

        idle_delta = idle - previous[0]
        total_delta = total - previous[1]

        if total_delta <= 0:
            return None

        return round(100.0 * (1.0 - idle_delta / total_delta), 2)

    @staticmethod
    def memory_percent() -> float | None:
        try:
            values: dict[str, int] = {}
            with open("/proc/meminfo") as handle:
                for line in handle:
                    key, _, rest = line.partition(":")
                    values[key] = int(rest.split()[0])
        except (OSError, ValueError, IndexError):
            return None

        total = values.get("MemTotal", 0)
        if total <= 0:
            return None

        # MemAvailable accounts for reclaimable cache; MemFree alone reports
        # almost every healthy Linux box as nearly out of memory.
        available = values.get("MemAvailable", values.get("MemFree", 0))

        return round(100.0 * (total - available) / total, 2)

    @staticmethod
    def disk_percent(path: str = "/") -> float | None:
        try:
            stat = os.statvfs(path)
        except OSError:
            return None

        total = stat.f_blocks * stat.f_frsize
        if total <= 0:
            return None

        free = stat.f_bavail * stat.f_frsize

        return round(100.0 * (total - free) / total, 2)

    @staticmethod
    def uptime_seconds() -> int | None:
        try:
            with open("/proc/uptime") as handle:
                return int(float(handle.readline().split()[0]))
        except (OSError, ValueError, IndexError):
            return None

    @staticmethod
    def interface_bytes(interface: str) -> tuple[int, int]:
        """Cumulative (rx, tx) for an interface, from sysfs."""
        base = f"/sys/class/net/{interface}/statistics"
        try:
            with open(f"{base}/rx_bytes") as handle:
                rx = int(handle.read().strip())
            with open(f"{base}/tx_bytes") as handle:
                tx = int(handle.read().strip())
            return rx, tx
        except (OSError, ValueError):
            return 0, 0


# ------------------------------------------------------------------- WireGuard


@dataclass
class Peer:
    public_key: str
    allowed_ips: str
    preshared_key: str | None = None

    def identity(self) -> tuple[str, str, str]:
        """Comparable form. Two peers are the same only if all three match."""
        return (self.public_key, self.allowed_ips, self.preshared_key or "")


class WireGuard:
    """Thin wrapper over `wg` and `ip`."""

    def __init__(self, interface: str) -> None:
        self.interface = interface

    def exists(self) -> bool:
        return os.path.isdir(f"/sys/class/net/{self.interface}")

    def ensure_interface(self, address: str | None, listen_port: int, mtu: int) -> None:
        """Creates and configures the interface if it is not already up.

        The node's own private key is generated here and never leaves the
        machine — the control plane stores only the public key it reports.
        """
        if not self.exists():
            log.info("Creating interface %s", self.interface)
            run(["ip", "link", "add", "dev", self.interface, "type", "wireguard"])

        key_path = f"/etc/wireguard/{self.interface}.key"

        if not os.path.exists(key_path):
            log.info("Generating node private key at %s", key_path)
            os.makedirs("/etc/wireguard", mode=0o700, exist_ok=True)

            private_key = run(["wg", "genkey"]).strip()

            # Written 0600 before any content lands in it.
            fd = os.open(key_path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
            with os.fdopen(fd, "w") as handle:
                handle.write(private_key + "\n")

        run(["wg", "set", self.interface,
             "listen-port", str(listen_port),
             "private-key", key_path])

        if address:
            current = run(["ip", "-o", "addr", "show", "dev", self.interface], check=False)
            if address.split("/")[0] not in current:
                log.info("Assigning %s to %s", address, self.interface)
                run(["ip", "addr", "flush", "dev", self.interface], check=False)
                run(["ip", "addr", "add", address, "dev", self.interface])

        if mtu:
            run(["ip", "link", "set", "mtu", str(mtu), "dev", self.interface], check=False)

        run(["ip", "link", "set", self.interface, "up"], check=False)

    def public_key(self) -> str | None:
        try:
            return run(["wg", "show", self.interface, "public-key"]).strip() or None
        except RuntimeError:
            return None

    def current_peers(self) -> dict[str, Peer]:
        """Peers currently configured, keyed by public key."""
        peers: dict[str, Peer] = {}

        try:
            dump = run(["wg", "show", self.interface, "dump"])
        except RuntimeError:
            return peers

        # First line describes the interface itself; peers follow.
        for line in dump.strip().splitlines()[1:]:
            fields = line.split("\t")
            if len(fields) < 4:
                continue

            public_key, preshared, _endpoint, allowed_ips = fields[:4]

            peers[public_key] = Peer(
                public_key=public_key,
                allowed_ips=allowed_ips,
                # `wg` prints a literal "(none)" rather than an empty field.
                preshared_key=None if preshared in ("(none)", "") else preshared,
            )

        return peers

    def transfer(self) -> list[dict[str, Any]]:
        """Per-peer counters and last handshake, from `wg show dump`."""
        entries: list[dict[str, Any]] = []

        try:
            dump = run(["wg", "show", self.interface, "dump"])
        except RuntimeError:
            return entries

        for line in dump.strip().splitlines()[1:]:
            fields = line.split("\t")
            if len(fields) < 7:
                continue

            try:
                entries.append({
                    "public_key": fields[0],
                    "last_handshake": int(fields[4]),
                    # WireGuard reports rx from the node's perspective. The
                    # panel stores it the same way, so no swap here.
                    "rx_bytes": int(fields[5]),
                    "tx_bytes": int(fields[6]),
                })
            except ValueError:
                continue

        return entries

    def add_peer(self, peer: Peer) -> None:
        args = ["wg", "set", self.interface,
                "peer", peer.public_key,
                "allowed-ips", peer.allowed_ips]

        if peer.preshared_key:
            # `wg` reads a preshared key from a file, never from argv, so it
            # cannot appear in the process list.
            with tempfile.NamedTemporaryFile("w", delete=False) as handle:
                os.chmod(handle.name, 0o600)
                handle.write(peer.preshared_key + "\n")
                psk_path = handle.name
            try:
                run(args + ["preshared-key", psk_path])
            finally:
                os.unlink(psk_path)
        else:
            run(args)

    def remove_peer(self, public_key: str) -> None:
        run(["wg", "set", self.interface, "peer", public_key, "remove"])


# ------------------------------------------------------------------ API client


class PanelClient:
    def __init__(self, config: Config) -> None:
        self.config = config
        self.base = f"{config.panel_url}/api/v1/node"

    def _request(self, method: str, path: str, payload: dict | None = None) -> dict:
        url = f"{self.base}{path}"
        data = json.dumps(payload).encode() if payload is not None else None

        request = urllib.request.Request(url, data=data, method=method)
        request.add_header("Authorization", f"Bearer {self.config.token}")
        request.add_header("Accept", "application/json")
        request.add_header("User-Agent", USER_AGENT)
        if data:
            request.add_header("Content-Type", "application/json")

        context = None
        if not self.config.verify_tls:
            import ssl
            context = ssl._create_unverified_context()

        with urllib.request.urlopen(request, timeout=self.config.timeout, context=context) as response:
            body = response.read().decode()
            return json.loads(body) if body else {}

    def heartbeat(self, payload: dict) -> dict:
        return self._request("POST", "/heartbeat", payload)

    def state(self) -> dict:
        return self._request("GET", "/state")

    def acknowledge(self, state_hash: str) -> dict:
        return self._request("POST", "/state/ack", {"state_hash": state_hash})

    def usage(self, peers: list[dict]) -> dict:
        return self._request("POST", "/usage", {"peers": peers})


# ----------------------------------------------------------------------- agent


@dataclass
class Agent:
    config: Config
    client: PanelClient
    wg: WireGuard
    stats: SystemStats = field(default_factory=SystemStats)
    running: bool = True
    applied_hash: str | None = None
    _last_usage: float = 0.0

    def stop(self, *_: Any) -> None:
        log.info("Shutdown requested; finishing current cycle.")
        self.running = False

    def run_forever(self) -> None:
        signal.signal(signal.SIGTERM, self.stop)
        signal.signal(signal.SIGINT, self.stop)

        log.info("zylo-agent %s starting (panel=%s interface=%s)",
                 VERSION, self.config.panel_url, self.config.interface)

        backoff = self.config.poll_interval

        while self.running:
            try:
                self.tick()
                # Reset the backoff only after a clean cycle.
                backoff = self.config.poll_interval
            except urllib.error.HTTPError as error:
                if error.code == 401:
                    # A wrong token will never fix itself by retrying faster.
                    log.error("Authentication rejected (401). Check the agent token.")
                    backoff = min(backoff * 2, 300)
                else:
                    log.error("Panel returned HTTP %s: %s", error.code, error.reason)
                    backoff = min(backoff * 2, 120)
            except (urllib.error.URLError, socket.timeout, TimeoutError) as error:
                # Transient: the panel may be restarting or the link may be down.
                # Existing tunnels keep working regardless — the data plane does
                # not depend on the control plane being reachable.
                log.warning("Panel unreachable (%s). Tunnels are unaffected.", error)
                backoff = min(backoff * 2, 120)
            except Exception as error:  # noqa: BLE001 - the loop must not die
                log.exception("Unexpected error: %s", error)
                backoff = min(backoff * 2, 120)

            for _ in range(backoff):
                if not self.running:
                    break
                time.sleep(1)

        log.info("zylo-agent stopped.")

    def tick(self) -> None:
        response = self.client.heartbeat(self.build_heartbeat())

        desired_hash = response.get("state_hash")

        if desired_hash and desired_hash != self.applied_hash:
            log.info("State changed (%s -> %s); syncing.",
                     (self.applied_hash or "none")[:12], desired_hash[:12])
            self.sync()

        if time.time() - self._last_usage >= self.config.usage_interval:
            self.report_usage()
            self._last_usage = time.time()

    def build_heartbeat(self) -> dict:
        rx, tx = SystemStats.interface_bytes(self.config.interface)

        payload = {
            "cpu_usage": self.stats.cpu_percent(),
            "ram_usage": SystemStats.memory_percent(),
            "disk_usage": SystemStats.disk_percent(),
            "rx_bytes": rx,
            "tx_bytes": tx,
            "active_peers": len(self.wg.current_peers()),
            "uptime_seconds": SystemStats.uptime_seconds(),
            "wireguard_running": self.wg.exists(),
            "applied_state_hash": self.applied_hash,
        }

        # Reported so the panel can record it on first contact without an
        # operator pasting it by hand. The panel accepts it only while unset.
        public_key = self.wg.public_key()
        if public_key:
            payload["public_key"] = public_key

        return {k: v for k, v in payload.items() if v is not None}

    def sync(self) -> None:
        """Fetches desired state and reconciles the interface toward it."""
        state = self.client.state()

        server = state.get("server", {})
        address = server.get("address")
        listen_port = int(server.get("listen_port") or 51820)
        mtu = int(server.get("mtu") or 1420)

        if not address:
            # No IP pool configured yet. Bringing the interface up without an
            # address would produce a tunnel that accepts handshakes and then
            # blackholes every packet, which is worse than staying down.
            log.warning("Panel reports no interface address; node is not provisioned yet.")
            return

        self.wg.ensure_interface(address, listen_port, mtu)

        desired = {
            p["public_key"]: Peer(
                public_key=p["public_key"],
                allowed_ips=p["allowed_ips"],
                preshared_key=p.get("preshared_key"),
            )
            for p in state.get("peers", [])
        }
        current = self.wg.current_peers()

        added = removed = updated = 0

        # Remove first, so an address being reassigned from one peer to another
        # never has two claimants on the interface at the same instant.
        for public_key in current.keys() - desired.keys():
            self.wg.remove_peer(public_key)
            removed += 1

        for public_key, peer in desired.items():
            existing = current.get(public_key)

            if existing is None:
                self.wg.add_peer(peer)
                added += 1
            elif existing.identity() != peer.identity():
                # `wg set` on an existing peer updates it in place.
                self.wg.add_peer(peer)
                updated += 1

        log.info("Sync complete: +%d ~%d -%d (%d peers total)",
                 added, updated, removed, len(desired))

        state_hash = state.get("state_hash")
        if state_hash:
            self.client.acknowledge(state_hash)
            self.applied_hash = state_hash

    def report_usage(self) -> None:
        entries = self.wg.transfer()

        if not entries:
            return

        result = self.client.usage(entries)
        log.debug("Usage reported: %s accepted, %s ignored",
                  result.get("accepted"), result.get("ignored"))


# ------------------------------------------------------------------------ main


def main() -> int:
    config_path = os.environ.get("ZYLO_AGENT_CONFIG", DEFAULT_CONFIG)

    if "--version" in sys.argv:
        print(f"zylo-agent {VERSION}")
        return 0

    logging.basicConfig(
        level=logging.DEBUG if "--debug" in sys.argv else logging.INFO,
        # No timestamp: journald adds its own, and duplicates make logs harder
        # to read, not easier.
        format="%(levelname)s %(message)s",
        stream=sys.stdout,
    )

    if os.geteuid() != 0:
        log.error("zylo-agent must run as root to configure WireGuard.")
        return 1

    require_binaries()

    config = Config.load(config_path)
    agent = Agent(
        config=config,
        client=PanelClient(config),
        wg=WireGuard(config.interface),
    )

    if "--once" in sys.argv:
        # Single cycle, for smoke-testing an install before enabling the unit.
        agent.tick()
        return 0

    agent.run_forever()
    return 0


if __name__ == "__main__":
    sys.exit(main())
