#!/usr/bin/python3
"""ovpmon-helper: the only root-side entry point for the unprivileged ovpmon services.

Usage (via doas):  ovpmon-helper install-config
                   ovpmon-helper service start|stop|restart|status
install-config reads the staged OpenVPN server config, validates it against a strict
allowlist of directives and installs it atomically to /etc/openvpn/server.conf.
"""
import ipaddress
import json
import os
import pwd
import re
import shlex
import stat
import subprocess
import sys

STAGED = "/var/lib/ovpmon/staging/server.conf"
TARGET = "/etc/openvpn/server.conf"
PKI_DIR = "/opt/OpenVPN-Monitoring-Simple/APP_PROFILER/easy-rsa/pki"
SCRIPTS_DIR = "/etc/openvpn/scripts"
STATUS_LOG = "/var/log/openvpn/openvpn-status.log"
SERVICE_USER = "ovpmon"
MAX_SIZE = 64 * 1024
CIPHERS_RE = re.compile(r"^[A-Za-z0-9:_-]{1,200}$")
os.environ["PATH"] = "/usr/sbin:/usr/bin:/sbin:/bin"


class Reject(Exception):
    pass


def under(path, base):
    real = os.path.realpath(path)
    return real == base or real.startswith(base.rstrip("/") + "/")


def need(cond, msg):
    if not cond:
        raise Reject(msg)


def is_int(x, lo, hi):
    return re.fullmatch(r"\d{1,6}", x) is not None and lo <= int(x) <= hi


def valid_route(r):
    parts = r.split()
    try:
        if len(parts) == 1:
            ipaddress.IPv4Network(parts[0], strict=False)
        elif len(parts) == 2:
            ipaddress.IPv4Address(parts[0])
            ipaddress.IPv4Network("0.0.0.0/" + parts[1])
        else:
            return False
        return True
    except ValueError:
        return False


def check_script(path):
    need(re.fullmatch(r"/etc/openvpn/scripts/[A-Za-z0-9_.-]{1,64}", path), "script path not allowed")
    real = os.path.realpath(path)
    need(os.path.dirname(real) == SCRIPTS_DIR and os.path.isfile(real), "script must be a file in " + SCRIPTS_DIR)
    st, dst = os.stat(real), os.stat(SCRIPTS_DIR)
    need(st.st_uid == 0 and not st.st_mode & 0o022, "script must be root-owned and not group/other-writable")
    need(dst.st_uid == 0 and not dst.st_mode & 0o022, SCRIPTS_DIR + " must be root-owned and not writable")


def check_line(tokens):
    d, a = tokens[0], tokens[1:]
    if d == "dev":
        need(a == ["tun"], "dev must be tun")
    elif d == "proto":
        need(len(a) == 1 and a[0] in ("udp", "tcp", "udp4", "tcp4", "udp6", "tcp6"), "bad proto")
    elif d in ("tls-server", "client-to-client", "duplicate-cn", "persist-key", "persist-tun"):
        need(not a, d + " takes no arguments")
    elif d == "explicit-exit-notify":
        need(len(a) == 1 and is_int(a[0], 1, 10), "bad explicit-exit-notify")
    elif d in ("port", "management-port"):
        need(len(a) == 1 and is_int(a[0], 1, 65535), "bad port")
    elif d in ("ca", "cert", "key", "dh", "crl-verify"):
        need(len(a) == 1 and under(a[0], PKI_DIR), d + " must be a file inside the PKI directory")
    elif d == "tls-auth":
        need(len(a) == 2 and under(a[0], PKI_DIR) and a[1] in ("0", "1"), "bad tls-auth")
    elif d == "tun-mtu":
        need(len(a) == 1 and is_int(a[0], 576, 9000), "bad tun-mtu")
    elif d == "mssfix":
        need(len(a) == 1 and is_int(a[0], 536, 1500), "bad mssfix")
    elif d == "topology":
        need(a == ["subnet"], "topology must be subnet")
    elif d == "server":
        need(len(a) == 2, "bad server")
        net = ipaddress.IPv4Network(f"{a[0]}/{a[1]}", strict=True)
        need(8 <= net.prefixlen <= 30, "bad server prefix")
    elif d == "ifconfig-pool-persist":
        need(a == ["/etc/openvpn/ipp.txt"], "ifconfig-pool-persist path not allowed")
    elif d in ("log", "log-append"):
        need(a == ["/var/log/openvpn/openvpn.log"], d + " path not allowed")
    elif d == "verb":
        need(len(a) == 1 and is_int(a[0], 0, 9), "bad verb")
    elif d == "status":
        need(len(a) == 2 and a[0] == STATUS_LOG and is_int(a[1], 1, 3600), "bad status")
    elif d == "status-version":
        need(a in (["1"], ["2"], ["3"]), "bad status-version")
    elif d == "push":
        need(len(a) == 1, "push takes one quoted argument")
        p = a[0]
        if p == "redirect-gateway def1 bypass-dhcp":
            return
        m = re.fullmatch(r"route (.+)", p)
        if m:
            need(valid_route(m.group(1)), "bad pushed route")
            return
        m = re.fullmatch(r"dhcp-option DNS (\S+)", p)
        need(m is not None, "pushed option not allowed")
        ipaddress.ip_address(m.group(1))
    elif d == "user":
        need(a == ["nobody"], "user must be nobody")
    elif d == "group":
        need(a == ["nogroup"], "group must be nogroup")
    elif d in ("data-ciphers", "data-ciphers-fallback"):
        need(len(a) == 1 and CIPHERS_RE.match(a[0]), "bad cipher list")
    elif d == "auth":
        need(len(a) == 1 and a[0] in ("SHA256", "SHA384", "SHA512"), "bad auth")
    elif d == "keepalive":
        need(len(a) == 2 and is_int(a[0], 1, 3600) and is_int(a[1], 1, 7200), "bad keepalive")
    elif d == "script-security":
        need(a == ["2"], "script-security must be 2")
    elif d in ("client-connect", "client-disconnect"):
        need(len(a) == 1, d + " takes one argument")
        check_script(a[0])
    elif d == "management":
        need(len(a) == 2 and ipaddress.ip_address(a[0]).is_loopback and is_int(a[1], 1, 65535), "management must be loopback")
    else:
        raise Reject("directive not allowed: " + d)


def validate(text):
    seen = set()
    for n, raw in enumerate(text.splitlines(), 1):
        line = raw.strip()
        if not line or line.startswith("#") or line.startswith(";"):
            continue
        need(all(32 <= ord(c) < 127 for c in line), f"line {n}: non-printable or non-ASCII character")
        try:
            tokens = shlex.split(line, comments=False)
        except ValueError as e:
            raise Reject(f"line {n}: {e}")
        try:
            check_line(tokens)
        except Reject as e:
            raise Reject(f"line {n}: {e}")
        except ValueError as e:
            raise Reject(f"line {n}: invalid value ({e})")
        seen.add(tokens[0])
    for req in ("user", "group", "server", "ca", "cert", "key"):
        need(req in seen, f"required directive missing: {req}")


def install_config():
    uid = pwd.getpwnam(SERVICE_USER).pw_uid
    fd = os.open(STAGED, os.O_RDONLY | os.O_NOFOLLOW)
    try:
        st = os.fstat(fd)
        need(st.st_uid == uid and stat.S_ISREG(st.st_mode), "staged config must be a regular file owned by " + SERVICE_USER)
        need(st.st_size <= MAX_SIZE, "staged config too large")
        data = os.read(fd, MAX_SIZE + 1)
    finally:
        os.close(fd)
    text = data.decode("ascii")  # one read: validate exactly what gets installed
    validate(text)
    tmp = TARGET + ".tmp"
    fd = os.open(tmp, os.O_WRONLY | os.O_CREAT | os.O_TRUNC | os.O_NOFOLLOW, 0o644)
    with os.fdopen(fd, "w") as f:
        f.write(text)
    os.chmod(tmp, 0o644)
    os.replace(tmp, TARGET)
    if not os.path.lexists("/etc/openvpn/openvpn.conf"):
        os.symlink("server.conf", "/etc/openvpn/openvpn.conf")


def prepare_status_log():
    """Let the unprivileged monitoring gatherer read the status log."""
    import grp
    gid = grp.getgrnam(SERVICE_USER).gr_gid
    if not os.path.exists(STATUS_LOG):
        open(STATUS_LOG, "a").close()
    os.chown(STATUS_LOG, 0, gid)
    os.chmod(STATUS_LOG, 0o640)


def service(action):
    need(action in ("start", "stop", "restart", "status"), "invalid action")
    if action in ("start", "restart"):
        need(os.path.isfile(TARGET), "server.conf is not installed")
        prepare_status_log()
    r = subprocess.run(["/sbin/rc-service", "openvpn", action], capture_output=True, text=True, timeout=60)
    return r.returncode, (r.stdout + r.stderr).strip()[-500:]


def main(argv):
    try:
        if argv == ["install-config"]:
            install_config()
            print(json.dumps({"status": "ok"}))
            return 0
        if len(argv) == 2 and argv[0] == "service":
            rc, out = service(argv[1])
            print(json.dumps({"status": "ok" if rc == 0 else "error", "output": out}))
            return 0 if rc == 0 else 2
        raise Reject("usage: install-config | service start|stop|restart|status")
    except Reject as e:
        print(json.dumps({"status": "rejected", "error": str(e)}))
        return 3
    except Exception as e:  # never leak a traceback with paths to the caller
        print(json.dumps({"status": "error", "error": type(e).__name__ + ": " + str(e)[:200]}))
        return 4


if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
