#!/usr/bin/python3 -I
"""Narrow, root-owned nftables helper. Never accepts rules, paths or shell code."""
import fcntl
import hashlib
import http.client
import ipaddress
import json
import os
from pathlib import Path
import re
import socket
import ssl
import stat
import subprocess
import sys
from concurrent.futures import ThreadPoolExecutor

ROOT = Path('/var/lib/zond-guard')
STATE = ROOT / 'state.json'
MARK = 0x5A6F6E64
NFT = '/usr/sbin/nft'
IP = '/usr/sbin/ip'


def transport_routes(remove=False):
    # File capabilities make /proc/PID/exe unreadable to the TUN classifier.
    # Route only Zond's privileged mark through the normal routing table;
    # an absent physical route must fail, never loop back into the TUN.
    for family in ('-4', '-6'):
        def command(*args):
            p = subprocess.run([IP, family, *args], text=True, capture_output=True,
                               timeout=5, env={'PATH': '/usr/sbin:/usr/bin:/sbin:/bin', 'LC_ALL': 'C'})
            if p.returncode:
                raise RuntimeError('Transport routing failed')
            return p.stdout
        rules = json.loads(command('-j', 'rule', 'show'))
        for priority, table in ((8998, 'main'), (8999, None)):
            matches = [r for r in rules if r.get('priority') == priority
                       and int(str(r.get('fwmark', '0')), 0) == MARK
                       and int(str(r.get('fwmask', '0xffffffff')), 0) == 0xffffffff
                       and (r.get('table') in ('main', 254) if table else r.get('action') == 'unreachable')]
            args = ['priority', str(priority), 'fwmark', hex(MARK)]
            args += ['lookup', table] if table else ['type', 'unreachable']
            if remove:
                for _ in matches:
                    command('rule', 'del', *args)
            elif not matches:
                command('rule', 'add', *args)


def trusted(path):
    for p in [path, *path.parents]:
        s = p.lstat()
        if stat.S_ISLNK(s.st_mode) or s.st_uid != 0 or s.st_mode & 0o022:
            raise ValueError('Untrusted protection path')


def nft(*args, data=None, optional=False):
    p = subprocess.run([NFT, *args], input=data, text=True, capture_output=True,
                       timeout=8, env={'PATH': '/usr/sbin:/usr/bin:/sbin:/bin', 'LC_ALL': 'C'})
    if p.returncode and not optional:
        raise RuntimeError('The kernel rejected protection rules')
    return p


def snapshot():
    # Listing all tables distinguishes an absent table from a permissions/kernel error.
    tables = json.loads(nft('-j', 'list', 'tables').stdout)['nftables']
    if not any(x.get('table', {}).get('name') == 'zond_guard' and
               x.get('table', {}).get('family') == 'inet' for x in tables):
        return None
    def stable(x):
        if isinstance(x, dict):
            return {k: stable(v) for k, v in x.items() if k not in ('handle', 'metainfo')}
        if isinstance(x, list):
            return [stable(v) for v in x if not isinstance(v, dict) or 'metainfo' not in v]
        return x
    data = stable(json.loads(nft('-j', 'list', 'table', 'inet', 'zond_guard').stdout))
    return hashlib.sha256(json.dumps(data, sort_keys=True).encode()).hexdigest()


def load():
    if not STATE.exists():
        return {'enabled': False}
    trusted(STATE)
    data = json.loads(STATE.read_text())
    if type(data.get('enabled')) is not bool:
        raise ValueError('Invalid protection state')
    return data


def save(data):
    temporary = ROOT / 'state.tmp'
    fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_TRUNC | os.O_NOFOLLOW, 0o600)
    with os.fdopen(fd, 'w') as f:
        json.dump(data, f)
        f.flush()
        os.fsync(f.fileno())
    os.replace(temporary, STATE)
    fd = os.open(ROOT, os.O_DIRECTORY)
    try:
        os.fsync(fd)
    finally:
        os.close(fd)


def install():
    old = snapshot()
    # A single transaction: a failed replacement leaves the previous filter intact.
    rules = ('delete table inet zond_guard\n' if old else '') + '''
table inet zond_guard {
 chain output {
  type filter hook output priority -10; policy drop;
  oifname "lo" accept
  meta mark 0x5a6f6e64 accept
  oifname "Lodnet" accept
  udp sport 68 udp dport 67 accept
  ip6 daddr ff02::1:2 udp sport 546 udp dport 547 accept
  ip6 hoplimit 255 icmpv6 type { nd-router-solicit, nd-neighbor-solicit, nd-neighbor-advert } accept
 }
 chain forward {
  type filter hook forward priority -10; policy drop;
 }
}
'''
    nft('-f', '-', data=rules)
    state = load()
    state.update(enabled=True, fingerprint=snapshot())
    save(state)


class Bootstrap(http.client.HTTPSConnection):
    def __init__(self, address, hostname):
        super().__init__(hostname, timeout=4)
        self.address = address

    def connect(self):
        # No system resolver, redirects, proxy variables or user-supplied URLs.
        sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        sock.settimeout(4)
        try:
            sock.setsockopt(socket.SOL_SOCKET, socket.SO_MARK, MARK)
            sock.connect((self.address, 443))
            self.sock = ssl.create_default_context().wrap_socket(sock, server_hostname=self.host)
        except BaseException:
            sock.close()
            raise


def domain(host):
    if not isinstance(host, str):
        raise ValueError('Invalid VPN host')
    host = host.lower().rstrip('.')
    if len(host) > 253 or not re.fullmatch(r'[a-zA-Z0-9](?:[a-zA-Z0-9.-]*[a-zA-Z0-9])?', host):
        raise ValueError('Invalid VPN host')
    if any(not label or len(label) > 63 or label.startswith('-') or label.endswith('-') for label in host.split('.')):
        raise ValueError('Invalid VPN host')
    return host


def resolve(host):
    host = domain(host)
    state = load()
    # A separate, passwordless action can resolve only peers authorized when
    # protection was enabled. It cannot change the list, firewall or resolver URL.
    if not state['enabled'] or host not in state.get('peers', []):
        raise PermissionError('VPN peer was not authorized')
    def query(record):
        for address, hostname, path in (('1.1.1.1', 'cloudflare-dns.com', '/dns-query'),
                                         ('8.8.8.8', 'dns.google', '/resolve')):
            connection = Bootstrap(address, hostname)
            try:
                connection.request('GET', path + '?name=' + host + '&type=' + record,
                                   headers={'Accept': 'application/dns-json'})
                response = connection.getresponse()
                body = response.read(32769)
                if response.status != 200 or len(body) > 32768:
                    raise ValueError('Invalid bootstrap response')
                document = json.loads(body)
                if document.get('Status') in (0, 3):
                    results = []
                    for answer in document.get('Answer', []):
                        if answer.get('type') == (1 if record == 'A' else 28):
                            results.append(str(ipaddress.ip_address(answer['data'])))
                    return results
            except (OSError, ValueError, http.client.HTTPException):
                pass
            finally:
                connection.close()
        return []
    with ThreadPoolExecutor(max_workers=2) as pool:
        results = [ip for row in pool.map(query, ('A', 'AAAA')) for ip in row]
    if not results:
        raise RuntimeError('VPN bootstrap DNS returned no addresses')
    print(json.dumps(list(dict.fromkeys(results))[:32]))


def main():
    if os.geteuid() != 0:
        raise PermissionError('Administrator authorization required')
    trusted(Path(__file__).absolute())
    command = sys.argv[1:]
    if command[:1] == ['resolve'] and len(command) == 2:
        resolve(command[1])
        return
    if len(command) != 1 or command[0] not in ('status', 'enable', 'disable', 'boot', 'prepare', 'uninstall'):
        raise ValueError('Unknown protection action')
    ROOT.mkdir(mode=0o700, exist_ok=True)
    trusted(ROOT)
    fd = os.open(ROOT / 'lock', os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW, 0o600)
    with os.fdopen(fd, 'w') as lock:
        fcntl.flock(lock, fcntl.LOCK_EX)
        state = load()
        action = command[0]
        if action == 'status':
            current = snapshot()
            print('On' if state['enabled'] and current and current == state.get('fingerprint') else
                  'Off' if not state['enabled'] and current is None else 'Partial')
        elif action == 'enable':
            if not Path('/sys/class/net/Lodnet').exists():
                raise ValueError('Connect Zond TUN before enabling protection')
            data = sys.stdin.read(32769)
            if len(data) > 32768:
                raise ValueError('Too many VPN peers')
            peers = json.loads(data or '[]')
            if not isinstance(peers, list) or len(peers) > 256:
                raise ValueError('Too many VPN peers')
            peers = sorted(set(domain(peer) for peer in peers))
            save({'enabled': True, 'peers': peers})
            transport_routes()
            install()
        elif action in ('boot', 'prepare'):
            # Installed even when protection is off: marked Xray must bypass TUN.
            transport_routes()
            if action == 'boot' and state['enabled']:
                install()
        else:
            if snapshot() is not None:
                nft('delete', 'table', 'inet', 'zond_guard')
            save({'enabled': False})
            if action == 'uninstall':
                transport_routes(remove=True)


if __name__ == '__main__':
    try:
        main()
    except Exception as error:
        # Never include provider input, responses or command output in logs.
        print('Protection operation failed: ' + type(error).__name__, file=sys.stderr)
        sys.exit(1)
