#!/usr/bin/python3

# SPDX-FileCopyrightText: (c) 2026 Daniel Hast
#
# SPDX-License-Identifier: Apache-2.0 OR MIT

import os
import sys
import tomllib
from dataclasses import dataclass
from pathlib import Path
from typing import Final

CONFIG_FILE: Final[Path] = Path("/etc/bwrap-restricted/config.toml")
DEFAULT_BWRAP_PATH: Final[Path] = Path("/usr/bin/bwrap")

# Mapping from bwrap options to the number of arguments each option takes
BWRAP_OPTIONS: Final[dict[bytes, int]] = {
    b"--help": 0,
    b"--version": 0,
    b"--args": 1,
    b"--argv0": 1,
    b"--level-prefix": 0,
    b"--unshare-user": 0,
    b"--unshare-user-try": 0,
    b"--unshare-ipc": 0,
    b"--unshare-pid": 0,
    b"--unshare-net": 0,
    b"--unshare-uts": 0,
    b"--unshare-cgroup": 0,
    b"--unshare-cgroup-try": 0,
    b"--unshare-all": 0,
    b"--share-net": 0,
    b"--userns": 1,
    b"--userns2": 1,
    b"--disable-userns": 0,
    b"--assert-userns-disabled": 0,
    b"--pidns": 1,
    b"--uid": 1,
    b"--gid": 1,
    b"--hostname": 1,
    b"--chdir": 1,
    b"--setenv": 2,
    b"--unsetenv": 1,
    b"--clearenv": 0,
    b"--lock-file": 1,
    b"--sync-fd": 1,
    b"--perms": 1,
    b"--size": 1,
    b"--bind": 2,
    b"--bind-try": 2,
    b"--dev-bind": 2,
    b"--dev-bind-try": 2,
    b"--ro-bind": 2,
    b"--ro-bind-try": 2,
    b"--remount-ro": 1,
    b"--overlay-src": 1,
    b"--overlay": 3,
    b"--tmp-overlay": 1,
    b"--ro-overlay": 1,
    b"--proc": 1,
    b"--dev": 1,
    b"--tmpfs": 1,
    b"--mqueue": 1,
    b"--dir": 1,
    b"--file": 2,
    b"--bind-data": 2,
    b"--ro-bind-data": 2,
    b"--symlink": 2,
    b"--chmod": 2,
    b"--seccomp": 1,
    b"--add-seccomp-fd": 1,
    b"--exec-label": 1,
    b"--file-label": 1,
    b"--block-fd": 1,
    b"--userns-block-fd": 1,
    b"--info-fd": 1,
    b"--json-status-fd": 1,
    b"--new-session": 0,
    b"--die-with-parent": 0,
    b"--as-pid-1": 0,
    b"--cap-add": 1,
    b"--cap-drop": 1,
}


def to_list_of_bytes(items: list[str]) -> list[bytes]:
    """Converts a list of strings to a list of byte-arrays"""
    return [bytes(s, encoding="utf8") for s in items]


@dataclass(frozen=True)
class Config:
    bwrap_path: Path
    forbidden_options: list[bytes]
    extra_options: list[bytes]

    @classmethod
    def from_file(cls, path: Path) -> "Config":
        with open(path, "rb") as f:
            config_data = tomllib.load(f)
        bwrap_path = Path(config_data.get("bwrap_path", DEFAULT_BWRAP_PATH))
        forbidden_options = config_data.get("forbidden_options", [])
        extra_options = config_data.get("extra_options", [])
        return cls(
            bwrap_path=bwrap_path,
            forbidden_options=to_list_of_bytes(forbidden_options),
            extra_options=to_list_of_bytes(extra_options),
        )

    @classmethod
    def default(cls) -> "Config":
        return cls(
            bwrap_path=DEFAULT_BWRAP_PATH,
            forbidden_options=[],
            extra_options=[],
        )


class BwrapOptionError(Exception):
    """The provided bwrap options are not accepted."""


def scan_bwrap_option(
    argv: list[bytes], i: int, config: Config, *, in_file: bool = False
) -> tuple[int, bool]:
    """
    Continue scanning bwrap options, starting at the given index.
    Returns a tuple of:
    - An integer indicating how many arguments were consumed;
    - A boolean indicating whether we're done parsing options.
    """
    argc: Final[int] = len(argv)
    arg = argv[i]
    arg_str = arg.decode()
    if arg == b"--" or not arg.startswith(b"-"):
        return 0, True
    if arg not in BWRAP_OPTIONS:
        raise BwrapOptionError(f"Unknown bwrap option: {arg_str}")
    option_argc = BWRAP_OPTIONS[arg]
    remaining_argc = argc - 1 - i
    if arg in config.forbidden_options:
        raise BwrapOptionError(f"Forbidden bwrap option: {arg_str}")
    if remaining_argc < option_argc:
        raise BwrapOptionError(
            f"Too few arguments after option {arg_str} (needed {option_argc}, found {remaining_argc})"
        )
    if arg == b"--args":
        if in_file:
            raise BwrapOptionError("--args option not accepted inside file descriptor")
        scan_args_fd(argv[i + 1], config)
        return 2, False
    return 1 + option_argc, False


def scan_args_fd(fd_arg: bytes, config: Config) -> None:
    """Scan bwrap arguments in file descriptor"""
    fd_argv = []
    chunk_size = 4096
    try:
        fd = int(fd_arg)
        buf = b""
        while True:
            buf += os.read(fd, chunk_size)
            if not buf:
                break
            new_args = buf.split(b"\0")
            fd_argv += new_args[:-1]
            buf = new_args[-1]
        os.lseek(fd, 0, os.SEEK_SET)
    except (ValueError, OSError) as err:
        raise BwrapOptionError(f"Bad file descriptor: --args {fd_arg.decode()}") from err

    i = 0
    while i < len(fd_argv):
        try:
            consumed, options_done = scan_bwrap_option(fd_argv, i, config, in_file=True)
        except BwrapOptionError as err:
            raise BwrapOptionError(f"{err} (in --args {fd})") from err
        i += consumed
        if options_done:
            break
    if i < len(fd_argv):
        raise BwrapOptionError(f"Non-option arguments in file descriptor (in --args {fd})")


def parse_bwrap_args(argv: list[bytes], config: Config) -> list[bytes]:
    """
    Parse bwrap args and return the modified arguments to be run, or raise a
    BwrapOptionError if parsing fails or a forbidden option is present.
    """
    i = 1
    while i < len(argv):
        consumed, options_done = scan_bwrap_option(argv, i, config)
        i += consumed
        if options_done:
            break
    return [bytes(config.bwrap_path), *argv[1:i], *config.extra_options, *argv[i:]]


def print_err(msg: str) -> None:
    print(f"bwrap-restricted: {msg}", file=sys.stderr)


def run_with_config(argv: list[bytes], config: Config) -> int:
    """Run with a given argument list and config"""
    try:
        bwrap_argv = parse_bwrap_args(argv, config)
    except BwrapOptionError as err:
        print_err(f"Error: {err}")
        return 1

    try:
        return os.execv(bwrap_argv[0], bwrap_argv)
    except OSError as err:
        print_err(f"Failed to execute bwrap: {err}")
        return 1


def main(argv: list[bytes] | None = None, config_file: Path = CONFIG_FILE) -> int:
    if argv is None:
        argv = [bytes(arg, encoding="utf8") for arg in sys.argv]
    try:
        config = Config.from_file(config_file)
    except FileNotFoundError:
        config = Config.default()
    except OSError:
        print_err(f"failed to open config file {config_file}")
        return 1
    except ValueError as err:
        print_err(str(err))
        return 1

    return run_with_config(argv, config)


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