"""Ownership checks for paths the agent executes, or hands to a root-run
service, on a caller's behalf.

Whoever can write such a file, or rename a component of its path, runs
code as root (DEF-56083). So the file and every directory above it must
belong to root and be writable by root only.
"""
import os
import stat

# Single syscall entry points. Unit tests may run unprivileged, so they
# substitute these to present the files they create as root-owned.
_stat = os.stat
_lstat = os.lstat

_WRITABLE_BY_OTHERS = stat.S_IWGRP | stat.S_IWOTH


def _check_directory(path, directory, st):
    if st.st_uid != 0:
        raise ValueError(
            "{}: directory {} is not owned by root (uid {})".format(
                path, directory, st.st_uid
            )
        )
    if st.st_mode & _WRITABLE_BY_OTHERS:
        raise ValueError(
            "{}: directory {} is group- or world-writable".format(
                path, directory
            )
        )


_MAX_SYMLINK_HOPS = 40  # the kernel's own limit


def _check_lexical_chain(path):
    """Every component of ``path`` as spelled, links included, must be
    root-owned: a link the caller owns can be repointed after the check
    however trustworthy its current target is. Directories must also not
    be writable by others; a symlink's own mode bits carry no meaning.

    The walk follows the kernel's rules rather than normalising the
    string: a root-owned symlink is resolved one hop at a time and its
    target's components are walked and checked like any other, so a
    later ``..`` steps out of the link's target, not out of the link's
    parent directory, and nothing the kernel traverses goes unchecked.
    """
    current = os.sep if os.path.isabs(path) else os.getcwd()
    pending = [p for p in path.split(os.sep) if p and p != "."]
    hops = 0
    while pending:
        part = pending.pop(0)
        if part == "..":
            current = os.path.dirname(current) or os.sep
            continue
        here = os.path.join(current, part)
        try:
            st = _lstat(here)
        except OSError as exc:
            raise ValueError("{}: cannot stat {}: {}".format(path, here, exc))
        if stat.S_ISLNK(st.st_mode):
            if st.st_uid != 0:
                raise ValueError(
                    "{}: symlink {} is not owned by root (uid {})".format(
                        path, here, st.st_uid
                    )
                )
            hops += 1
            if hops > _MAX_SYMLINK_HOPS:
                raise ValueError("{}: too many symlinks".format(path))
            target = os.readlink(here)
            if os.path.isabs(target):
                current = os.sep
            pending = [
                p for p in target.split(os.sep) if p and p != "."
            ] + pending
        else:
            if stat.S_ISDIR(st.st_mode) and pending:
                _check_directory(path, here, st)
            current = here


def verify_root_owned_path(path):
    """Return ``realpath(path)`` if only root could have written it.

    ``path`` must resolve to a regular file owned by uid 0 that is not
    group- or world-writable; every directory above the resolved file up
    to ``/`` must be owned by uid 0 and not group- or world-writable; and
    every component of ``path`` as spelled, symlinks included, must be
    owned by uid 0 as well, since the spelling is what gets stored and
    re-read later. A sticky bit does not exempt a writable directory.

    Raise :class:`ValueError` naming the first offending component.
    Callers that go on to execute or persist the file should use the
    returned real path rather than ``path``.
    """
    real = os.path.realpath(path)
    try:
        st = _stat(real)
    except FileNotFoundError:
        raise ValueError("{} does not exist".format(path))
    except OSError as exc:
        raise ValueError("{}: stat failed: {}".format(path, exc))
    if not stat.S_ISREG(st.st_mode):
        raise ValueError("{} is not a regular file".format(path))
    if st.st_uid != 0:
        raise ValueError(
            "{} is not owned by root (uid {})".format(path, st.st_uid)
        )
    if st.st_mode & _WRITABLE_BY_OTHERS:
        raise ValueError("{} is group- or world-writable".format(path))

    directory = os.path.dirname(real)
    while True:
        try:
            dst = _stat(directory)
        except OSError as exc:
            raise ValueError(
                "{}: cannot stat directory {}: {}".format(path, directory, exc)
            )
        _check_directory(path, directory, dst)
        if directory == "/":
            break
        directory = os.path.dirname(directory)

    _check_lexical_chain(path)
    return real
