#!/usr/bin/env python3
import argparse
import ctypes
import errno
import os
from pathlib import Path
import stat
import uuid
RENAME_NOREPLACE = 1
LIBC = ctypes.CDLL(None, use_errno=True)


def rename_noreplace(descriptor, source, destination):
    result = LIBC.renameat2(
        descriptor,
        os.fsencode(source),
        descriptor,
        os.fsencode(destination),
        RENAME_NOREPLACE,
    )
    if result != 0:
        error = ctypes.get_errno()
        raise OSError(error, os.strerror(error))


def drop_privileges(uid, gid):
    if os.geteuid() == uid and os.getegid() == gid:
        return
    if os.geteuid() != 0:
        raise RuntimeError("cannot assume target identity")
    os.setgroups([])
    os.setgid(gid)
    os.setuid(uid)
    if os.geteuid() != uid or os.getegid() != gid:
        raise RuntimeError("target identity transition failed")


def relative_parts(value):
    path = Path(value)
    if path.is_absolute() or not path.parts or any(part in {".", ".."} for part in path.parts):
        raise RuntimeError("invalid target-home relative path")
    return path.parts


def require_owner(descriptor, uid):
    if os.fstat(descriptor).st_uid != uid:
        raise RuntimeError("target-home directory has unexpected owner")


def open_home(home, uid):
    if not home.is_absolute() or home == Path("/"):
        raise RuntimeError("invalid target home")
    descriptor = os.open(home, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
    try:
        require_owner(descriptor, uid)
        return descriptor
    except Exception:
        os.close(descriptor)
        raise


def open_directory(root, parts, uid):
    descriptor = os.dup(root)
    try:
        for part in parts:
            child = os.open(
                part,
                os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
                dir_fd=descriptor,
            )
            require_owner(child, uid)
            os.close(descriptor)
            descriptor = child
        return descriptor
    except Exception:
        os.close(descriptor)
        raise


def ensure_directory(home, relative, uid, gid, mode):
    root = open_home(home, uid)
    descriptor = os.dup(root)
    try:
        for part in relative_parts(relative):
            try:
                child = os.open(
                    part,
                    os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
                    dir_fd=descriptor,
                )
                require_owner(child, uid)
            except FileNotFoundError:
                os.mkdir(part, mode=mode, dir_fd=descriptor)
                os.fsync(descriptor)
                child = os.open(
                    part,
                    os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
                    dir_fd=descriptor,
                )
                require_owner(child, uid)
                os.fchmod(child, mode)
                os.fsync(child)
            os.close(descriptor)
            descriptor = child
    finally:
        os.close(descriptor)
        os.close(root)


def parent_binding(home, relative, uid):
    parts = relative_parts(relative)
    root = open_home(home, uid)
    try:
        parent = open_directory(root, parts[:-1], uid)
    finally:
        os.close(root)
    return parent, parts[-1]


def replace_symlink(home, relative, target, uid, gid):
    parent, leaf = parent_binding(home, relative, uid)
    temporary = f".{leaf}.install-{uuid.uuid4().hex}"
    try:
        os.symlink(target, temporary, dir_fd=parent)
        os.chown(temporary, uid, gid, dir_fd=parent, follow_symlinks=False)
        try:
            rename_noreplace(parent, temporary, leaf)
        except OSError as error:
            if error.errno != errno.EEXIST:
                raise
            info = os.stat(leaf, dir_fd=parent, follow_symlinks=False)
            if (
                not stat.S_ISLNK(info.st_mode)
                or os.readlink(leaf, dir_fd=parent) != target
                or info.st_uid != uid
                or info.st_gid != gid
            ):
                raise RuntimeError("managed target-home path contains an unrelated object") from error
            os.unlink(temporary, dir_fd=parent)
        os.fsync(parent)
    finally:
        try:
            os.unlink(temporary, dir_fd=parent)
        except FileNotFoundError:
            pass
        os.close(parent)


def unlink_leaf(home, relative, uid, gid, expected_target=None):
    parent, leaf = parent_binding(home, relative, uid)
    quarantine = f".{leaf}.remove-{uuid.uuid4().hex}"
    try:
        try:
            os.rename(leaf, quarantine, src_dir_fd=parent, dst_dir_fd=parent)
        except FileNotFoundError:
            return
        os.fsync(parent)
        info = os.stat(quarantine, dir_fd=parent, follow_symlinks=False)
        valid = stat.S_ISLNK(info.st_mode) and info.st_uid == uid and info.st_gid == gid
        if valid and expected_target is not None:
            valid = os.readlink(quarantine, dir_fd=parent) == expected_target
        if not valid:
            os.rename(quarantine, leaf, src_dir_fd=parent, dst_dir_fd=parent)
            os.fsync(parent)
            raise RuntimeError("refusing to unlink unexpected target-home object")
        os.unlink(quarantine, dir_fd=parent)
        os.fsync(parent)
    finally:
        try:
            os.rename(quarantine, leaf, src_dir_fd=parent, dst_dir_fd=parent)
            os.fsync(parent)
        except FileNotFoundError:
            pass
        os.close(parent)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("action", choices=("ensure-directory", "symlink", "unlink"))
    parser.add_argument("--home", required=True, type=Path)
    parser.add_argument("--relative", required=True)
    parser.add_argument("--uid", type=int)
    parser.add_argument("--gid", type=int)
    parser.add_argument("--mode", type=lambda value: int(value, 8), default=0o755)
    parser.add_argument("--target")
    args = parser.parse_args()
    if args.uid is None or args.uid < 0:
        parser.error("uid is required")
    if args.action in {"ensure-directory", "symlink"} and (args.gid is None or args.gid < 0):
        parser.error("gid is required")
    if args.gid is None:
        parser.error("gid is required")
    drop_privileges(args.uid, args.gid)
    if args.action == "ensure-directory":
        ensure_directory(args.home, args.relative, args.uid, args.gid, args.mode)
    elif args.action == "symlink":
        if args.target is None or not Path(args.target).is_absolute():
            parser.error("absolute symlink target is required")
        replace_symlink(args.home, args.relative, args.target, args.uid, args.gid)
    else:
        unlink_leaf(args.home, args.relative, args.uid, args.gid, args.target)


if __name__ == "__main__":
    main()
