#!/usr/bin/python -I
"""
Composant privilégié de Meshamoto — monte et démonte le tunnel WireGuard.

Invoqué par le client via pkexec, sous une règle polkit `allow_active` : tout utilisateur
d'une session locale active peut l'appeler SANS mot de passe. C'est tout l'intérêt du
composant, et c'est aussi ce qui rend son contrat critique.

CONTRAINTE STRUCTURANTE — ne jamais accepter de fichier de configuration.
Un `.conf` WireGuard accepte des directives `PostUp = <commande shell>` que `wg-quick`
exécute EN ROOT. Un helper sans mot de passe qui accepterait un fichier écrit par la
session offrirait donc un root gratuit à tout code tournant sous ce compte. Ce helper
reçoit uniquement des PARAMÈTRES, les valide un à un contre des motifs stricts, et
FABRIQUE lui-même le fichier à un emplacement que l'utilisateur ne peut pas écrire.

Les motifs de validation refusent tout saut de ligne : sans ça, une « clé publique »
contenant `\\nPostUp = ...` réintroduirait exactement l'injection qu'on cherche à
empêcher.

La clé privée arrive sur l'ENTRÉE STANDARD, jamais en argument : `/proc/<pid>/cmdline`
est lisible par n'importe quel utilisateur de la machine.

`#!/usr/bin/python -I` (mode isolé) est délibéré : il neutralise PYTHONPATH,
PYTHONSTARTUP et le site-packages de l'utilisateur, qui feraient sinon charger du code
arbitraire EN ROOT. Ne pas retirer cet indicateur, et n'importer que la bibliothèque
standard — ce script ne doit dépendre d'aucun chemin que l'utilisateur peut écrire.
"""

from __future__ import annotations

import argparse
import ipaddress
import os
import re
import shutil
import subprocess
import sys

TUNNEL_NAME = "meshamoto"

# Répertoire du .conf fabriqué : sur tmpfs, appartenant à root, mode 0700. L'utilisateur
# ne doit pas pouvoir remplacer le fichier entre sa fabrication et sa lecture par
# wg-quick — c'est pourquoi on n'écrit PAS dans le dossier de configuration du client.
RUNTIME_DIR = "/run/meshamoto"

# Clé WireGuard : 32 octets en base64, soit 43 caractères + '='.
#
# \Z et NON $ : en Python, `$` accepte un saut de ligne FINAL, donc "<cle>=\n" passerait
# la validation. Le principe « aucun saut de ligne, jamais » doit être absolu, c'est lui
# qui empêche l'injection de directives PostUp dans le fichier fabriqué.
KEY_RE = re.compile(r"\A[A-Za-z0-9+/]{43}=\Z")


def _fail(message: str) -> "typing.NoReturn":  # noqa: F821
    print(f"meshamoto-helper: {message}", file=sys.stderr)
    raise SystemExit(2)


def _valid_key(value: str, what: str) -> str:
    if not KEY_RE.match(value):
        # Ne JAMAIS réafficher la valeur reçue : pour la clé privée, le message
        # d'erreur finirait dans le journal du client, lisible par la session.
        _fail(f"{what} invalide (attendu : 44 caracteres base64)")
    return value


def _valid_interface_address(value: str) -> str:
    """
    Adresse locale du tunnel, sous la forme <ipv4>/<prefixe>.

    Le préfixe est EXIGÉ explicitement : `ipaddress.ip_interface("192.168.192.2")` est
    accepté sans erreur et vaut /32, ce qui donnerait une interface incapable de router
    vers ses pairs. Un défaut silencieux est pire qu'un refus.
    """
    if "/" not in value:
        _fail(f"prefixe manquant dans l'adresse d'interface : {value!r}")
    try:
        iface = ipaddress.ip_interface(value)
    except ValueError:
        _fail(f"adresse d'interface invalide : {value!r}")
    if iface.version != 4:
        _fail("seul IPv4 est pris en charge")
    return str(iface)


def _valid_allowed_ip(value: str) -> str:
    try:
        return str(ipaddress.ip_network(value, strict=True))
    except ValueError:
        _fail(f"AllowedIPs invalide : {value!r}")


def _valid_port(value: str) -> int:
    try:
        port = int(value)
    except ValueError:
        _fail(f"port invalide : {value!r}")
    if not 1 <= port <= 65535:
        _fail(f"port hors plage : {port}")
    return port


def _valid_endpoint(value: str) -> str:
    """`<ipv4>:<port>` uniquement — pas de nom d'hote, pour eviter toute resolution
    declenchee en root a partir d'une valeur fournie par la session."""
    host, _, port = value.rpartition(":")
    if not host or not port:
        _fail(f"endpoint invalide : {value!r}")
    try:
        ipaddress.IPv4Address(host)
    except ValueError:
        _fail(f"endpoint invalide (IPv4 attendue) : {value!r}")
    _valid_port(port)
    return f"{host}:{port}"


def _parse_peer(raw: str) -> dict:
    """`<cle_publique>,<allowed_ip>[,<endpoint>]`"""
    parts = raw.split(",")
    if not 2 <= len(parts) <= 3:
        _fail(f"pair mal forme : {raw!r}")
    peer = {
        "public_key": _valid_key(parts[0], "cle publique"),
        "allowed_ip": _valid_allowed_ip(parts[1]),
        "endpoint": _valid_endpoint(parts[2]) if len(parts) == 3 and parts[2] else None,
    }
    return peer


def _render_conf(private_key: str, address: str, listen_port: int, peers: list[dict]) -> str:
    """
    Fabrique le .conf a partir de valeurs DEJA validees. Aucune donnee ne transite ici
    sans etre passee par les fonctions _valid_* : c'est ce qui garantit qu'aucune
    directive PostUp/PreUp ne peut etre injectee.
    """
    lines = [
        "[Interface]",
        f"PrivateKey = {private_key}",
        f"Address = {address}",
        f"ListenPort = {listen_port}",
        "",
    ]
    for peer in peers:
        lines += ["[Peer]", f"PublicKey = {peer['public_key']}", f"AllowedIPs = {peer['allowed_ip']}"]
        if peer["endpoint"]:
            lines.append(f"Endpoint = {peer['endpoint']}")
        lines += ["PersistentKeepalive = 25", ""]
    return "\n".join(lines)


def _render_sync_conf(private_key: str, listen_port: int, peers: list) -> str:
    """
    Variante depouillee pour `wg syncconf` : PAS de ligne Address.

    `Address` appartient a wg-quick et configure l'adresse et les routes ; `wg`
    ne la connait pas et echouerait dessus. C'est aussi ce qui fait l'interet de
    syncconf : il met a jour les pairs SANS demonter l'interface, donc sans la
    coupure que subissaient tous les pairs a chaque arrivee ou depart.

    Memes valeurs deja validees que _render_conf : aucune donnee n'arrive ici
    sans etre passee par les fonctions _valid_*.
    """
    lines = [
        "[Interface]",
        f"PrivateKey = {private_key}",
        f"ListenPort = {listen_port}",
        "",
    ]
    for peer in peers:
        lines += ["[Peer]", f"PublicKey = {peer['public_key']}", f"AllowedIPs = {peer['allowed_ip']}"]
        if peer["endpoint"]:
            lines.append(f"Endpoint = {peer['endpoint']}")
        lines += ["PersistentKeepalive = 25", ""]
    return "\n".join(lines)


def _conf_path() -> str:
    os.makedirs(RUNTIME_DIR, mode=0o700, exist_ok=True)
    os.chmod(RUNTIME_DIR, 0o700)
    return os.path.join(RUNTIME_DIR, f"{TUNNEL_NAME}.conf")


def _wg() -> str:
    """Chemin de `wg` (distinct de wg-quick) : seul lui sait faire syncconf."""
    for chemin in ("/usr/bin/wg", "/bin/wg"):
        if os.path.exists(chemin):
            return chemin
    return "/usr/bin/wg"


def _wg_quick() -> str:
    path = shutil.which("wg-quick", path="/usr/bin:/usr/sbin:/bin:/sbin")
    if path is None:
        _fail("wg-quick introuvable (paquet wireguard-tools)")
    return path


def _run(*args: str) -> subprocess.CompletedProcess:
    return subprocess.run(list(args), capture_output=True, text=True, timeout=60)


def _interface_exists() -> bool:
    return _run("/usr/bin/ip", "link", "show", TUNNEL_NAME).returncode == 0


def cmd_set_tunnel(args: argparse.Namespace) -> int:
    private_key = _valid_key(sys.stdin.read().strip(), "cle privee")
    address = _valid_interface_address(args.address)
    listen_port = _valid_port(args.listen_port)
    peers = [_parse_peer(raw) for raw in args.peer]

    path = _conf_path()
    # 0600 des la creation : la fenetre pendant laquelle la cle privee serait lisible
    # ne doit pas exister, meme brievement.
    fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
    with os.fdopen(fd, "w") as handle:
        handle.write(_render_conf(private_key, address, listen_port, peers))

    if _interface_exists():
        _run(_wg_quick(), "down", path)

    result = _run(_wg_quick(), "up", path)
    if result.returncode != 0:
        print(result.stderr, file=sys.stderr)
        return 1
    return 0


def cmd_sync_tunnel(args: argparse.Namespace) -> int:
    """
    Met a jour les pairs A CHAUD, sans demonter l'interface.

    Echoue volontairement si l'interface n'existe pas : c'est a l'appelant de
    retomber sur set-tunnel. Faire les deux ici masquerait au client l'etat reel
    du tunnel, alors qu'il en a besoin pour son diagnostic de panne locale.
    """
    if not _interface_exists():
        print("interface absente : utiliser set-tunnel", file=sys.stderr)
        return 1

    private_key = _valid_key(sys.stdin.read().strip(), "cle privee")
    listen_port = _valid_port(args.listen_port)
    peers = [_parse_peer(raw) for raw in args.peer]

    # Fichier temporaire en 0600 des la creation, comme pour set-tunnel : la cle
    # privee ne doit jamais etre lisible, meme brievement.
    path = os.path.join(RUNTIME_DIR, f"{TUNNEL_NAME}-sync.conf")
    fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
    try:
        with os.fdopen(fd, "w") as handle:
            handle.write(_render_sync_conf(private_key, listen_port, peers))
        result = _run(_wg(), "syncconf", TUNNEL_NAME, path)
    finally:
        try:
            os.unlink(path)
        except OSError:
            pass
    if result.returncode != 0:
        print(result.stderr, file=sys.stderr)
        return 1
    return 0


def cmd_stop_tunnel(_args: argparse.Namespace) -> int:
    if not _interface_exists():
        return 0
    path = os.path.join(RUNTIME_DIR, f"{TUNNEL_NAME}.conf")
    target = path if os.path.exists(path) else TUNNEL_NAME
    result = _run(_wg_quick(), "down", target)
    if result.returncode != 0 and "is not a WireGuard interface" not in result.stderr:
        print(result.stderr, file=sys.stderr)
        return 1
    return 0


# --- Mode isolation (nftables) ---------------------------------------------

TABLE_NAME = "meshamoto"

# Nom d'hote du DNS interne : labels DNS + suffixe fixe. Le client assainit deja les
# pseudos (hosts_file_common._sanitize_label), mais c'est ICI que la garantie compte :
# un pseudo vient d'un autre pair via le serveur, donc d'une source non fiable, et
# /etc/hosts est ecrit en root. Un saut de ligne qui passerait permettrait de detourner
# n'importe quel domaine sur la machine.
HOSTNAME_RE = re.compile(r"\A[a-z0-9-]+\.meshamoto\Z")

HOSTS_PATH = "/etc/hosts"
HOSTS_BLOCK_BEGIN = "# BEGIN MESHAMOTO"
HOSTS_BLOCK_END = "# END MESHAMOTO"


def _valid_ipv4(value: str) -> str:
    try:
        return str(ipaddress.IPv4Address(value))
    except ValueError:
        _fail(f"adresse IPv4 invalide : {value!r}")


def _valid_cidr(value: str) -> str:
    try:
        return str(ipaddress.ip_network(value, strict=True))
    except ValueError:
        _fail(f"CIDR invalide : {value!r}")


def _parse_acl(raw: str) -> tuple:
    """`<ipv4>=<port>[,<port>...]`"""
    ip, sep, ports_raw = raw.partition("=")
    if not sep or not ports_raw:
        _fail(f"ACL mal formee : {raw!r}")
    ports = [_valid_port(p) for p in ports_raw.split(",")]
    return _valid_ipv4(ip), ports


def _build_ruleset(virtual_cidr: str, listen_port: int, acls: list) -> str:
    """
    Ruleset nftables, a partir de valeurs deja validees.

    DOIT rester identique a client/firewall_backend_linux._build_ruleset : les deux
    chemins (helper privilegie et repli pkexec) doivent produire exactement les memes
    regles, sinon le comportement du mode isolation dependrait de la presence du
    composant. Un test croise verifie cette egalite (tests/test_privileged_helper.py).
    """
    lines = [
        f"add table inet {TABLE_NAME}",
        f"delete table inet {TABLE_NAME}",
        f"table inet {TABLE_NAME} {{",
        "    chain input {",
        "        type filter hook input priority filter; policy accept;",
        f"        udp dport {listen_port} accept",
    ]
    for peer_ip, ports in acls:
        port_list = ", ".join(str(p) for p in ports)
        lines.append(f"        ip saddr {peer_ip} tcp dport {{ {port_list} }} accept")
        lines.append(f"        ip saddr {peer_ip} udp dport {{ {port_list} }} accept")
        lines.append(f"        ip saddr {peer_ip} drop")
    lines += [
        f"        ip saddr {virtual_cidr} accept",
        "        ct state established,related accept",
        "        iif lo accept",
        "        drop",
        "    }",
        "}",
    ]
    return "\n".join(lines) + "\n"


def _nft() -> str:
    path = shutil.which("nft", path="/usr/bin:/usr/sbin:/bin:/sbin")
    if path is None:
        _fail("nft introuvable (paquet nftables)")
    return path


def cmd_set_firewall(args: argparse.Namespace) -> int:
    virtual_cidr = _valid_cidr(args.virtual_cidr)
    listen_port = _valid_port(args.listen_port)
    acls = [_parse_acl(raw) for raw in args.acl]

    path = os.path.join(RUNTIME_DIR, "isolation.nft")
    os.makedirs(RUNTIME_DIR, mode=0o700, exist_ok=True)
    fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
    with os.fdopen(fd, "w") as handle:
        handle.write(_build_ruleset(virtual_cidr, listen_port, acls))

    result = _run(_nft(), "-f", path)
    if result.returncode != 0:
        print(result.stderr, file=sys.stderr)
        return 1
    return 0


def cmd_clear_firewall(_args: argparse.Namespace) -> int:
    result = _run(_nft(), "delete", "table", "inet", TABLE_NAME)
    # Table absente : rien a faire, ce n'est pas une erreur.
    if result.returncode != 0 and "No such file or directory" not in result.stderr:
        print(result.stderr, file=sys.stderr)
        return 1
    return 0


# --- DNS interne (/etc/hosts) ----------------------------------------------

def _valid_hostname(value: str) -> str:
    if not HOSTNAME_RE.match(value):
        _fail(f"nom d'hote invalide : {value!r}")
    return value


def _parse_entry(raw: str) -> tuple:
    """`<hostname>.meshamoto=<ipv4>`"""
    name, sep, ip = raw.partition("=")
    if not sep:
        _fail(f"entree DNS mal formee : {raw!r}")
    return _valid_hostname(name), _valid_ipv4(ip)


def _strip_block(content: str) -> str:
    """Retire le bloc Meshamoto, laisse le reste du fichier strictement intact."""
    kept, inside = [], False
    for line in content.splitlines():
        stripped = line.strip()
        if stripped == HOSTS_BLOCK_BEGIN:
            inside = True
            continue
        if stripped == HOSTS_BLOCK_END:
            inside = False
            continue
        if not inside:
            kept.append(line)
    return "\n".join(kept)


def _build_hosts(existing: str, entries: list) -> str:
    base = _strip_block(existing).rstrip("\n")
    if not entries:
        return base + "\n" if base else ""
    block = [HOSTS_BLOCK_BEGIN]
    for hostname, ip in sorted(entries):
        block.append(f"{ip}\t{hostname}")
    block.append(HOSTS_BLOCK_END)
    parts = [base] if base else []
    parts.append("\n".join(block))
    return "\n".join(parts) + "\n"


def _write_hosts(content: str) -> None:
    """Ecriture atomique : /etc/hosts corrompu rendrait la machine difficile a utiliser."""
    temp = HOSTS_PATH + ".meshamoto.tmp"
    fd = os.open(temp, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o644)
    with os.fdopen(fd, "w") as handle:
        handle.write(content)
    os.replace(temp, HOSTS_PATH)


def cmd_set_hosts(args: argparse.Namespace) -> int:
    entries = [_parse_entry(raw) for raw in args.entry]
    with open(HOSTS_PATH, encoding="utf-8") as handle:
        existing = handle.read()
    _write_hosts(_build_hosts(existing, entries))
    return 0


def cmd_clear_hosts(_args: argparse.Namespace) -> int:
    with open(HOSTS_PATH, encoding="utf-8") as handle:
        existing = handle.read()
    _write_hosts(_build_hosts(existing, []))
    return 0


def main() -> int:
    if os.geteuid() != 0:
        _fail("doit etre execute en root (via pkexec)")

    parser = argparse.ArgumentParser(prog="meshamoto-helper", add_help=True)
    sub = parser.add_subparsers(dest="command", required=True)

    up = sub.add_parser("set-tunnel", help="monte ou remonte le tunnel")
    up.add_argument("--address", required=True, help="<ipv4>/<prefixe> de l'interface")
    up.add_argument("--listen-port", required=True)
    up.add_argument("--peer", action="append", default=[],
                    help="<cle_publique>,<allowed_ip>[,<endpoint>] (repetable)")
    up.set_defaults(func=cmd_set_tunnel)

    sync = sub.add_parser("sync-tunnel", help="met a jour les pairs a chaud (wg syncconf)")
    sync.add_argument("--listen-port", required=True)
    sync.add_argument("--peer", action="append", default=[],
                      help="<cle_publique>,<allowed_ip>[,<endpoint>] (repetable)")
    sync.set_defaults(func=cmd_sync_tunnel)

    down = sub.add_parser("stop-tunnel", help="demonte le tunnel")
    down.set_defaults(func=cmd_stop_tunnel)

    fw = sub.add_parser("set-firewall", help="applique le mode isolation")
    fw.add_argument("--virtual-cidr", required=True)
    fw.add_argument("--listen-port", required=True)
    fw.add_argument("--acl", action="append", default=[],
                    help="<ipv4>=<port>[,<port>...] (repetable)")
    fw.set_defaults(func=cmd_set_firewall)

    nofw = sub.add_parser("clear-firewall", help="retire le mode isolation")
    nofw.set_defaults(func=cmd_clear_firewall)

    hosts = sub.add_parser("set-hosts", help="ecrit le bloc DNS interne")
    hosts.add_argument("--entry", action="append", default=[],
                       help="<nom>.meshamoto=<ipv4> (repetable)")
    hosts.set_defaults(func=cmd_set_hosts)

    nohosts = sub.add_parser("clear-hosts", help="retire le bloc DNS interne")
    nohosts.set_defaults(func=cmd_clear_hosts)

    args = parser.parse_args()
    return args.func(args)


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