#!/usr/bin/python3
"""Root-owned IoTSploit host-state daemon. Standard library only."""

from __future__ import annotations

import argparse
import datetime
import grp
import hashlib
import ipaddress
import json
import os
import re
import selectors
import signal
import socket
import stat
import struct
import subprocess
import sys
import time
from pathlib import Path
from typing import Any


MAX_REQUEST_BYTES = 4_096
MAX_OUTPUT_BYTES = 8_192
MAX_RESPONSE_BYTES = 24_576
COMMAND_TIMEOUT_SECONDS = 10
SOCKET_BACKLOG = 16
DEFAULT_SOCKET_PATH = Path("/run/iotsploit/priv.sock")
IP_EXECUTABLE = "/usr/sbin/ip"
NMCLI_EXECUTABLE = "/usr/bin/nmcli"

CAN_INTERFACE = re.compile(r"^(v?can)[0-9]{1,3}$")
NETWORK_INTERFACE = re.compile(r"^[a-z0-9._-]{1,15}$")
VERB_SCHEMAS = {
    "can-fd-up": {
        "iface": "can",
        "bitrate": "integer",
        "sample_point": "ratio",
        "dbitrate": "integer",
        "dsample_point": "ratio",
    },
    "can-link-state": {"iface": "can", "state": ["up", "down"]},
    "can-up": {"iface": "can", "bitrate": "integer-or-null"},
    "doip-config": {"iface": "network"},
    "route-via": {"action": ["add", "delete"], "cidr": "ipv4-/16", "gateway": "ipv4"},
    "vlan-add": {
        "parent": "network",
        "vlan_id": "1-4094",
        "address": "ipv4-interface",
        "local_mac": "mac-or-null",
        "peer_ip": "ipv4-or-null",
        "peer_mac": "mac-or-null",
    },
    "vlan-edit": {
        "parent": "network",
        "vlan_id": "1-4094",
        "address": "ipv4-interface",
        "local_mac": "mac-or-null",
        "peer_ip": "ipv4-or-null",
        "peer_mac": "mac-or-null",
    },
    "vlan-delete": {"parent": "network", "vlan_id": "1-4094"},
}
VERB_KEYS = {verb: set(schema) for verb, schema in VERB_SCHEMAS.items()}
VERB_TABLE_HASH = hashlib.sha256(
    json.dumps(VERB_SCHEMAS, sort_keys=True, separators=(",", ":")).encode()
).hexdigest()


class RequestError(ValueError):
    pass


def _unique_object(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
    result: dict[str, Any] = {}
    for key, value in pairs:
        if key in result:
            raise RequestError(f"duplicate key: {key}")
        result[key] = value
    return result


def _read_request(connection: socket.socket) -> dict[str, Any]:
    data = bytearray()
    while b"\n" not in data:
        chunk = connection.recv(1_024)
        if not chunk:
            raise RequestError("request ended before newline")
        data.extend(chunk)
        if len(data) > MAX_REQUEST_BYTES:
            raise RequestError("request exceeds 4 KiB")
    line, trailing = bytes(data).split(b"\n", 1)
    while not trailing:
        chunk = connection.recv(1_024)
        if not chunk:
            break
        trailing += chunk
        if len(data) + len(trailing) > MAX_REQUEST_BYTES:
            raise RequestError("request exceeds 4 KiB")
    if trailing:
        raise RequestError("trailing request bytes are not allowed")
    try:
        payload = json.loads(line, object_pairs_hook=_unique_object)
    except UnicodeDecodeError as exc:
        raise RequestError("request is not UTF-8") from exc
    except json.JSONDecodeError as exc:
        raise RequestError("request is not valid JSON") from exc
    if not isinstance(payload, dict) or set(payload) != {"verb", "args"}:
        raise RequestError("request must contain exactly verb and args")
    if not isinstance(payload["verb"], str) or not isinstance(payload["args"], dict):
        raise RequestError("verb must be a string and args must be an object")
    return payload


def _text(value: Any, name: str) -> str:
    if not isinstance(value, str) or not value or "\x00" in value:
        raise RequestError(f"{name} must be a non-empty string without NUL")
    return value


def _can_interface(value: Any) -> str:
    interface = _text(value, "iface")
    if not CAN_INTERFACE.fullmatch(interface):
        raise RequestError("iface must match ^(v?can)[0-9]{1,3}$")
    return interface


def _network_interface(value: Any) -> str:
    interface = _text(value, "iface")
    if not NETWORK_INTERFACE.fullmatch(interface):
        raise RequestError("iface must match ^[a-z0-9._-]{1,15}$")
    return interface


def _bitrate(value: Any, name: str) -> int:
    if type(value) is not int or not 10_000 <= value <= 10_000_000:
        raise RequestError(f"{name} must be an integer from 10000 to 10000000")
    return value


def _sample_point(value: Any, name: str) -> str:
    """A CAN sample point as ip(8) wants it: a ratio, three decimals.

    Bounded well inside 0..1 because the ends are not sample points at all --
    a controller cannot sample a bit before it starts or after it ends.
    """
    if type(value) not in (int, float) or not 0.5 <= value <= 0.95:
        raise RequestError(f"{name} must be a number from 0.5 to 0.95")
    return f"{float(value):.3f}"


def _ipv4(value: Any, name: str) -> str:
    text = _text(value, name)
    try:
        address = ipaddress.ip_address(text)
    except ValueError as exc:
        raise RequestError(f"{name} must be an IPv4 address") from exc
    if address.version != 4:
        raise RequestError(f"{name} must be an IPv4 address")
    return str(address)


def _ipv4_interface(value: Any) -> str:
    text = _text(value, "address")
    if "/" not in text:
        raise RequestError("address must be an IPv4 interface in CIDR notation")
    try:
        interface = ipaddress.ip_interface(text)
    except ValueError as exc:
        raise RequestError("address must be an IPv4 interface in CIDR notation") from exc
    if interface.version != 4 or interface.ip.is_unspecified or interface.ip.is_multicast:
        raise RequestError("address must be a usable IPv4 interface in CIDR notation")
    return str(interface)


def _optional_ipv4(value: Any, name: str) -> str | None:
    if value is None:
        return None
    address = _ipv4(value, name)
    parsed = ipaddress.ip_address(address)
    if parsed.is_unspecified or parsed.is_multicast:
        raise RequestError(f"{name} must be a usable IPv4 address")
    return address


def _mac(value: Any, name: str) -> str:
    address = _text(value, name).lower()
    if not re.fullmatch(r"(?:[0-9a-f]{2}:){5}[0-9a-f]{2}", address):
        raise RequestError(f"{name} must be a colon-separated MAC address")
    octets = bytes.fromhex(address.replace(":", ""))
    if octets == b"\x00" * 6 or octets == b"\xff" * 6 or octets[0] & 1:
        raise RequestError(f"{name} must be a unicast MAC address")
    return address


def _optional_mac(value: Any, name: str) -> str | None:
    if value is None:
        return None
    return _mac(value, name)


def _vlan_id(value: Any) -> int:
    if type(value) is not int or not 1 <= value <= 4_094:
        raise RequestError("vlan_id must be an integer from 1 to 4094")
    return value


def _vlan_identity(args: dict[str, Any]) -> tuple[str, int, str, str]:
    parent = _network_interface(args["parent"])
    vlan_id = _vlan_id(args["vlan_id"])
    interface = f"{parent}.{vlan_id}"
    if len(interface) > 15:
        raise RequestError("derived VLAN interface name exceeds 15 characters")
    return parent, vlan_id, interface, f"iotsploit-vlan-{parent}-{vlan_id}"


def _vlan_config(args: dict[str, Any]) -> dict[str, Any]:
    parent, vlan_id, interface, profile = _vlan_identity(args)
    peer_ip = _optional_ipv4(args["peer_ip"], "peer_ip")
    peer_mac = _optional_mac(args["peer_mac"], "peer_mac")
    if (peer_ip is None) != (peer_mac is None):
        raise RequestError("peer_ip and peer_mac must both be set or both be null")
    return {
        "parent": parent,
        "vlan_id": vlan_id,
        "interface": interface,
        "profile": profile,
        "address": _ipv4_interface(args["address"]),
        "local_mac": _optional_mac(args["local_mac"], "local_mac"),
        "peer_ip": peer_ip,
        "peer_mac": peer_mac,
    }


def _neighbor_command(config: dict[str, Any]) -> list[str] | None:
    if config["peer_ip"] is None:
        return None
    return [
        IP_EXECUTABLE,
        "neigh",
        "replace",
        config["peer_ip"],
        "lladdr",
        config["peer_mac"],
        "dev",
        config["interface"],
        "nud",
        "permanent",
    ]


def _cidr(value: Any) -> str:
    text = _text(value, "cidr")
    try:
        network = ipaddress.ip_network(text, strict=False)
    except ValueError as exc:
        raise RequestError("cidr must be an IPv4 network") from exc
    if network.version != 4 or network.num_addresses > 65_536:
        raise RequestError("cidr must be IPv4 and no larger than /16")
    return str(network)


def _validate_request(payload: dict[str, Any]) -> tuple[str, dict[str, Any], list[list[str]]]:
    verb = payload["verb"]
    args = payload["args"]
    if verb not in VERB_KEYS:
        raise RequestError(f"unknown verb: {verb}")
    if set(args) != VERB_KEYS[verb]:
        raise RequestError(f"{verb} requires exactly: {', '.join(sorted(VERB_KEYS[verb]))}")

    if verb == "can-up":
        interface = _can_interface(args["iface"])
        bitrate = args["bitrate"]
        if interface.startswith("vcan"):
            if bitrate is not None:
                raise RequestError("vcan bitrate must be null")
            validated = {"iface": interface, "bitrate": None}
            commands = [[IP_EXECUTABLE, "link", "set", "dev", interface, "up"]]
        else:
            validated = {"iface": interface, "bitrate": _bitrate(bitrate, "physical CAN bitrate")}
            commands = [
                [IP_EXECUTABLE, "link", "set", "dev", interface, "type", "can", "bitrate", str(bitrate)],
                [IP_EXECUTABLE, "link", "set", "dev", interface, "up"],
            ]
    elif verb == "can-fd-up":
        interface = _can_interface(args["iface"])
        if interface.startswith("vcan"):
            raise RequestError("a virtual CAN interface has no bit timing")
        bitrate = _bitrate(args["bitrate"], "bitrate")
        dbitrate = _bitrate(args["dbitrate"], "dbitrate")
        sample_point = _sample_point(args["sample_point"], "sample_point")
        dsample_point = _sample_point(args["dsample_point"], "dsample_point")
        validated = {
            "iface": interface,
            "bitrate": bitrate,
            "sample_point": sample_point,
            "dbitrate": dbitrate,
            "dsample_point": dsample_point,
        }
        # Bit timing cannot be set on a running link, so the link is lowered
        # first. Both halves are one verb because a link left down between two
        # calls is a bus nobody is listening to and nobody was told about.
        commands = [
            [IP_EXECUTABLE, "link", "set", "dev", interface, "down"],
            [
                IP_EXECUTABLE, "link", "set", "dev", interface, "type", "can",
                "bitrate", str(bitrate), "sample-point", sample_point,
                "dbitrate", str(dbitrate), "dsample-point", dsample_point,
                "fd", "on",
            ],
            [IP_EXECUTABLE, "link", "set", "dev", interface, "up"],
        ]
    elif verb == "can-link-state":
        interface = _can_interface(args["iface"])
        state_value = _text(args["state"], "state")
        if state_value not in {"up", "down"}:
            raise RequestError("state must be up or down")
        validated = {"iface": interface, "state": state_value}
        commands = [[IP_EXECUTABLE, "link", "set", "dev", interface, state_value]]
    elif verb == "doip-config":
        interface = _network_interface(args["iface"])
        validated = {"iface": interface}
        commands = [
            [IP_EXECUTABLE, "address", "replace", "169.254.58.58/16", "dev", interface],
            [IP_EXECUTABLE, "route", "replace", "169.254.0.0/16", "dev", interface],
        ]
    elif verb == "route-via":
        action = _text(args["action"], "action")
        if action not in {"add", "delete"}:
            raise RequestError("action must be add or delete")
        network = _cidr(args["cidr"])
        gateway = _ipv4(args["gateway"], "gateway")
        validated = {"action": action, "cidr": network, "gateway": gateway}
        commands = [[IP_EXECUTABLE, "route", action, network, "via", gateway]]
    elif verb in {"vlan-add", "vlan-edit"}:
        config = _vlan_config(args)
        validated = {key: config[key] for key in args}
        if verb == "vlan-add":
            commands = [[
                NMCLI_EXECUTABLE,
                "--wait",
                "10",
                "connection",
                "add",
                "type",
                "vlan",
                "con-name",
                config["profile"],
                "ifname",
                config["interface"],
                "dev",
                config["parent"],
                "id",
                str(config["vlan_id"]),
                "ipv4.method",
                "manual",
                "ipv4.addresses",
                config["address"],
                "ipv4.never-default",
                "yes",
                "ipv6.method",
                "disabled",
                "connection.autoconnect",
                "yes",
            ]]
            if config["local_mac"] is not None:
                commands[0].extend(["802-3-ethernet.cloned-mac-address", config["local_mac"]])
        else:
            commands = [[
                NMCLI_EXECUTABLE,
                "--wait",
                "10",
                "connection",
                "modify",
                config["profile"],
                "ipv4.method",
                "manual",
                "ipv4.addresses",
                config["address"],
                "ipv4.gateway",
                "",
                "ipv4.never-default",
                "yes",
                "ipv6.method",
                "disabled",
                "connection.autoconnect",
                "yes",
                "802-3-ethernet.cloned-mac-address",
                config["local_mac"] or "",
            ]]
        commands.append([
            NMCLI_EXECUTABLE,
            "--wait",
            "10",
            "connection",
            "up",
            config["profile"],
        ])
        if verb == "vlan-edit":
            commands.append([
                IP_EXECUTABLE,
                "neigh",
                "flush",
                "dev",
                config["interface"],
                "nud",
                "permanent",
            ])
        neighbor = _neighbor_command(config)
        if neighbor is not None:
            commands.append(neighbor)
    else:
        parent, vlan_id, _, profile = _vlan_identity(args)
        validated = {"parent": parent, "vlan_id": vlan_id}
        commands = [[
            NMCLI_EXECUTABLE,
            "--wait",
            "10",
            "connection",
            "delete",
            profile,
        ]]
    return verb, validated, commands


def _append_capped(buffer: bytearray, chunk: bytes) -> bool:
    available = MAX_OUTPUT_BYTES - len(buffer)
    if available > 0:
        buffer.extend(chunk[:available])
    return len(chunk) > max(available, 0)


def _terminate_group(process: subprocess.Popen[bytes]) -> None:
    try:
        os.killpg(process.pid, signal.SIGTERM)
        process.wait(timeout=1)
    except subprocess.TimeoutExpired:
        os.killpg(process.pid, signal.SIGKILL)
        process.wait(timeout=1)
    except ProcessLookupError:
        pass


def _run_command(argv: list[str]) -> tuple[int, str, str, bool]:
    process = subprocess.Popen(
        argv,
        stdin=subprocess.DEVNULL,
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        env={},
        close_fds=True,
        start_new_session=True,
    )
    assert process.stdout is not None and process.stderr is not None
    selector = selectors.DefaultSelector()
    selector.register(process.stdout, selectors.EVENT_READ, "stdout")
    selector.register(process.stderr, selectors.EVENT_READ, "stderr")
    buffers = {"stdout": bytearray(), "stderr": bytearray()}
    truncated = False
    deadline = time.monotonic() + COMMAND_TIMEOUT_SECONDS
    timed_out = False

    while selector.get_map():
        remaining = deadline - time.monotonic()
        if remaining <= 0 and not timed_out:
            timed_out = True
            _terminate_group(process)
        for key, _ in selector.select(timeout=max(0.0, min(0.1, remaining)) if not timed_out else 0.1):
            chunk = os.read(key.fileobj.fileno(), 4_096)
            if chunk:
                truncated = _append_capped(buffers[key.data], chunk) or truncated
            else:
                selector.unregister(key.fileobj)
    selector.close()
    return_code = process.wait()
    if timed_out:
        return_code = 124
        truncated = _append_capped(buffers["stderr"], b"command timed out\n") or truncated
    return (
        return_code,
        buffers["stdout"].decode("utf-8", errors="replace"),
        buffers["stderr"].decode("utf-8", errors="replace"),
        truncated,
    )


def _execute(commands: list[list[str]]) -> tuple[int, str, str, bool]:
    stdout_parts = []
    stderr_parts = []
    truncated = False
    exit_code = 0
    for command in commands:
        exit_code, stdout, stderr, command_truncated = _run_command(command)
        stdout_parts.append(stdout)
        stderr_parts.append(stderr)
        truncated = truncated or command_truncated
        if exit_code != 0:
            break
    stdout_bytes = "".join(stdout_parts).encode()
    stderr_bytes = "".join(stderr_parts).encode()
    if len(stdout_bytes) > MAX_OUTPUT_BYTES or len(stderr_bytes) > MAX_OUTPUT_BYTES:
        truncated = True
    stdout = stdout_bytes[:MAX_OUTPUT_BYTES].decode("utf-8", errors="replace")
    stderr = stderr_bytes[:MAX_OUTPUT_BYTES].decode("utf-8", errors="replace")
    return exit_code, stdout, stderr, truncated


def _encoded_response(response: dict[str, Any]) -> bytes:
    encoded = json.dumps(response, separators=(",", ":"), ensure_ascii=False).encode() + b"\n"
    while len(encoded) > MAX_RESPONSE_BYTES and (response["stdout"] or response["stderr"]):
        field = "stdout" if len(response["stdout"]) >= len(response["stderr"]) else "stderr"
        response[field] = response[field][:-512]
        response["output_truncated"] = True
        encoded = json.dumps(response, separators=(",", ":"), ensure_ascii=False).encode() + b"\n"
    return encoded


def _peer(connection: socket.socket) -> tuple[int, int]:
    pid, uid, _ = struct.unpack("3i", connection.getsockopt(socket.SOL_SOCKET, socket.SO_PEERCRED, 12))
    return pid, uid


def _audit(event: str, *, pid: int, uid: int, verb: str, args: dict[str, Any], **fields: Any) -> None:
    record = {
        "timestamp": datetime.datetime.now(datetime.timezone.utc).isoformat(),
        "event": event,
        "peer_pid": pid,
        "peer_uid": uid,
        "verb": verb,
        "args": args,
        **fields,
    }
    print(json.dumps(record, sort_keys=True, separators=(",", ":")), file=sys.stderr, flush=True)


def _handle_connection(connection: socket.socket) -> None:
    started = time.monotonic()
    connection.settimeout(2)
    try:
        pid, uid = _peer(connection)
    except OSError:
        pid, uid = -1, -1
    verb = "invalid"
    args: dict[str, Any] = {}
    try:
        payload = _read_request(connection)
        verb, args, commands = _validate_request(payload)
        _audit("start", pid=pid, uid=uid, verb=verb, args=args)
        exit_code, stdout, stderr, truncated = _execute(commands)
        response: dict[str, Any] = {
            "ok": exit_code == 0,
            "exit": exit_code,
            "stdout": stdout,
            "stderr": stderr,
        }
        if truncated:
            response["output_truncated"] = True
    except (RequestError, OSError, ValueError) as exc:
        response = {"ok": False, "exit": 2, "stdout": "", "stderr": str(exc)}
    delivered = True
    try:
        connection.sendall(_encoded_response(response))
    except OSError:
        # The caller vanished before reading the reply. The command already ran,
        # so the outcome belongs in the audit rather than in a dead daemon.
        delivered = False
    duration_ms = round((time.monotonic() - started) * 1_000, 3)
    _audit(
        "finish",
        pid=pid,
        uid=uid,
        verb=verb,
        args=args,
        exit=response["exit"],
        duration_ms=duration_ms,
        delivered=delivered,
    )


def _systemd_socket() -> socket.socket | None:
    try:
        listen_pid = int(os.getenv("LISTEN_PID", "0"))
        listen_fds = int(os.getenv("LISTEN_FDS", "0"))
    except ValueError:
        return None
    if listen_pid != os.getpid() or listen_fds != 1:
        return None
    listener = socket.socket(fileno=3)
    if listener.family != socket.AF_UNIX:
        raise RuntimeError("systemd fd 3 is not an AF_UNIX socket")
    return listener


def _container_socket(path: Path, group_name: str) -> socket.socket:
    path.parent.mkdir(mode=0o755, parents=True, exist_ok=True)
    if path.exists() or path.is_symlink():
        metadata = path.lstat()
        if metadata.st_uid != 0 or not stat.S_ISSOCK(metadata.st_mode):
            raise RuntimeError(f"refusing to replace non-root or non-socket path: {path}")
        path.unlink()
    group_id = grp.getgrnam(group_name).gr_gid
    listener = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
    listener.bind(str(path))
    os.chown(path, 0, group_id)
    os.chmod(path, 0o660)
    listener.listen(SOCKET_BACKLOG)
    return listener


def _serve_once(listener: socket.socket) -> None:
    connection, _ = listener.accept()
    with connection:
        try:
            _handle_connection(connection)
        except Exception as exc:
            # One connection must never take the privileged daemon down with it.
            _audit("error", pid=-1, uid=-1, verb="unknown", args={}, error=repr(exc))


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--socket", type=Path, default=DEFAULT_SOCKET_PATH)
    parser.add_argument("--group", default="iotsploit")
    options = parser.parse_args()
    listener = _systemd_socket() or _container_socket(options.socket, options.group)
    while True:
        _serve_once(listener)


if __name__ == "__main__":
    raise SystemExit(main())
