#!/usr/bin/python

# omarchy:summary=Safely signal a process selected in Activity Monitor
# omarchy:group=system
# omarchy:args=<TERM|KILL|APP_TERM> <PID> <START_TICKS>

"""Signal a same-user process or app root while guarding against PID reuse."""

from __future__ import annotations

import argparse
from dataclasses import dataclass
import os
import select
import signal
import sys
from pathlib import Path


APP_TERM_GRACE_MS = 3000
PF_KTHREAD = 0x00200000
DEAD_STATES = frozenset(("Z", "X", "x"))
PROTECTED_NAMES = frozenset(
    (
        b"quickshell",
        b"hyprland",
        b"uwsm",
        b"systemd",
        b"systemd-executo",
        b"systemd-executor",
        b"(sd-pam)",
        b"dbus-broker",
        b"dbus-broker-lau",
        b"omarchy-hyprlan",
    )
)
PROC_ROOT = Path("/proc")


def fail(message: str, status: int = 1) -> None:
    print(message, file=sys.stderr)
    raise SystemExit(status)


class InspectionError(Exception):
    """A procfs identity record could not be read or validated."""


@dataclass(frozen=True)
class ProcessStat:
    pid: int
    name: bytes
    state: str
    parent_pid: int
    flags: int
    start_ticks: int


@dataclass(frozen=True)
class ProcessStatus:
    pid: int
    state: str
    user_ids: tuple[int, int, int, int]


def read_process_stat(proc_root: Path, pid: int) -> ProcessStat:
    try:
        raw = (proc_root / str(pid) / "stat").read_bytes()
    except FileNotFoundError:
        raise InspectionError(f"Process {pid} no longer exists") from None
    except OSError as error:
        raise InspectionError(f"Could not inspect process {pid}: {error}") from error

    command_start = raw.find(b"(")
    command_end = raw.rfind(b") ")
    if command_start < 1 or command_end <= command_start:
        raise InspectionError(f"Process {pid} returned an invalid stat record")

    try:
        record_pid = int(raw[:command_start].strip())
    except ValueError:
        raise InspectionError(f"Process {pid} returned an invalid PID") from None
    if record_pid != pid:
        raise InspectionError(f"Process {pid} returned a mismatched stat record")

    name = raw[command_start + 1 : command_end]

    # The remaining fields begin with field 3 (state), making starttime
    # (field 22) zero-based index 19. Splitting after the final ") " also
    # preserves process names containing spaces or parentheses.
    fields = raw[command_end + 2 :].split()
    if len(fields) <= 19:
        raise InspectionError(f"Process {pid} returned an incomplete stat record")

    try:
        state = fields[0].decode("ascii")
        parent_pid = int(fields[1])
        flags = int(fields[6])
        start_ticks = int(fields[19])
    except (UnicodeDecodeError, ValueError):
        raise InspectionError(f"Process {pid} returned malformed stat fields") from None

    if len(state) != 1 or not state.isalpha():
        raise InspectionError(f"Process {pid} returned an invalid state")
    if parent_pid < 0 or flags < 0:
        raise InspectionError(f"Process {pid} returned invalid process metadata")
    if start_ticks <= 0:
        raise InspectionError(f"Process {pid} returned an invalid start time")

    return ProcessStat(record_pid, name, state, parent_pid, flags, start_ticks)


def read_process_status(proc_root: Path, pid: int) -> ProcessStatus:
    try:
        lines = (proc_root / str(pid) / "status").read_text(
            encoding="ascii", errors="strict"
        ).splitlines()
    except FileNotFoundError:
        raise InspectionError(f"Process {pid} no longer exists") from None
    except (OSError, UnicodeError) as error:
        raise InspectionError(f"Could not inspect process {pid}: {error}") from error

    values: dict[str, str] = {}
    for line in lines:
        key, separator, value = line.partition(":")
        if separator and key in ("Pid", "State", "Uid"):
            values[key] = value.strip()

    if set(values) != {"Pid", "State", "Uid"}:
        raise InspectionError(f"Process {pid} returned an incomplete status record")

    try:
        status_pid = int(values["Pid"])
        user_ids_raw = values["Uid"].split()
        if len(user_ids_raw) != 4:
            raise ValueError
        user_ids = tuple(int(value) for value in user_ids_raw)
    except ValueError:
        raise InspectionError(f"Process {pid} returned malformed ownership data") from None

    state_fields = values["State"].split()
    state = state_fields[0] if state_fields else ""
    if status_pid != pid:
        raise InspectionError(f"Process {pid} returned a mismatched status record")
    if len(state) != 1 or not state.isalpha():
        raise InspectionError(f"Process {pid} returned an invalid status state")
    if any(user_id < 0 for user_id in user_ids):
        raise InspectionError(f"Process {pid} returned invalid ownership data")

    return ProcessStatus(status_pid, state, user_ids)  # type: ignore[arg-type]


def protected_processes(proc_root: Path) -> set[int]:
    """Return this helper and every currently observable ancestor."""
    protected: set[int] = set()
    current_pid = os.getpid()

    # A process cannot have a legitimate ancestry chain this deep. The bound
    # also makes malformed procfs data fail closed without looping forever.
    for _ in range(256):
        if current_pid <= 0 or current_pid in protected:
            break
        protected.add(current_pid)
        try:
            current_pid = read_process_stat(proc_root, current_pid).parent_pid
        except InspectionError:
            break

    # getppid() remains useful if the parent exited between walking self/stat
    # and reading its record.
    parent_pid = os.getppid()
    if parent_pid > 0:
        protected.add(parent_pid)
    return protected


def reject_unsafe_target(
    process_stat: ProcessStat, process_status: ProcessStatus, expected_start_ticks: int
) -> None:
    if process_stat.start_ticks != expected_start_ticks:
        fail(f"Process {process_stat.pid} changed before it could be signaled")
    if process_stat.flags & PF_KTHREAD:
        fail("Refusing to signal a kernel thread", 77)
    if process_stat.state in DEAD_STATES or process_status.state in DEAD_STATES:
        fail("Refusing to signal a zombie or dead process", 77)
    if process_stat.name.lower() in PROTECTED_NAMES:
        fail("Refusing to signal a protected desktop-session process", 77)

    current_uid = os.getuid()
    if any(user_id != current_uid for user_id in process_status.user_ids):
        fail("Refusing to signal a process owned by another user", 77)


def safe_app_ancestor(
    process_stat: ProcessStat,
    process_status: ProcessStatus,
    expected_name: bytes,
) -> bool:
    """Return whether a parent is safe to adopt as this app's signal target."""
    return (
        process_stat.name == expected_name
        and not process_stat.flags & PF_KTHREAD
        and process_stat.state not in DEAD_STATES
        and process_status.state not in DEAD_STATES
        and all(user_id == os.getuid() for user_id in process_status.user_ids)
    )


def app_signal_target(
    proc_root: Path,
    selected_stat: ProcessStat,
    selected_fd: int,
    protected: set[int],
) -> tuple[ProcessStat, int, list[int]]:
    """Resolve the highest safe, direct same-name ancestor of a selected process."""
    current_stat = selected_stat
    current_fd = selected_fd
    opened_fds = [selected_fd]
    seen = {selected_stat.pid}

    # The bound fails closed on malformed procfs ancestry without allowing an
    # unbounded walk. Each retained pidfd also prevents a validated identity
    # from silently becoming a recycled PID while the chain is resolved.
    for _ in range(256):
        parent_pid = current_stat.parent_pid
        if parent_pid <= 1 or parent_pid in protected or parent_pid in seen:
            break

        try:
            parent_fd = os.pidfd_open(parent_pid)
        except (ProcessLookupError, PermissionError, OSError):
            break

        try:
            parent_stat = read_process_stat(proc_root, parent_pid)
            parent_status = read_process_status(proc_root, parent_pid)
            current_check = read_process_stat(proc_root, current_stat.pid)
        except InspectionError:
            os.close(parent_fd)
            break

        # Recheck the direct edge after opening the parent pidfd. This avoids
        # promoting a process that ceased to be the selected app's parent
        # during inspection.
        if (
            current_check.start_ticks != current_stat.start_ticks
            or current_check.parent_pid != parent_pid
            or current_check.name != current_stat.name
            or not safe_app_ancestor(
                parent_stat, parent_status, selected_stat.name
            )
        ):
            os.close(parent_fd)
            break

        opened_fds.append(parent_fd)
        seen.add(parent_pid)
        current_stat = parent_stat
        current_fd = parent_fd

    return current_stat, current_fd, opened_fds


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description=(
            "Signal a same-user process or matching app root if its PID and "
            "start time still match."
        )
    )
    parser.add_argument("signal_name", choices=("TERM", "KILL", "APP_TERM"))
    parser.add_argument("pid", type=int)
    parser.add_argument("start_ticks", type=int)
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    if args.pid <= 1:
        fail("Refusing to signal a protected system process", 77)
    if args.start_ticks <= 0:
        fail("Process start time must be positive", 64)
    protected = protected_processes(PROC_ROOT)
    if args.pid in protected:
        fail("Refusing to signal Activity Monitor or its ancestor chain", 77)

    try:
        pid_fd = os.pidfd_open(args.pid)
    except ProcessLookupError:
        fail(f"Process {args.pid} no longer exists")
    except OSError as error:
        fail(f"Could not open process {args.pid}: {error}")

    opened_fds = [pid_fd]
    app_result = ""
    target_stat: ProcessStat | None = None
    target_fd = pid_fd
    try:
        try:
            process_stat = read_process_stat(PROC_ROOT, args.pid)
            process_status = read_process_status(PROC_ROOT, args.pid)
        except InspectionError as error:
            fail(str(error))

        reject_unsafe_target(process_stat, process_status, args.start_ticks)
        if args.signal_name == "APP_TERM":
            target_stat, target_fd, opened_fds = app_signal_target(
                PROC_ROOT, process_stat, pid_fd, protected
            )
            try:
                target_check = read_process_stat(PROC_ROOT, target_stat.pid)
                target_status = read_process_status(PROC_ROOT, target_stat.pid)
            except InspectionError as error:
                fail(str(error))
            reject_unsafe_target(
                target_check, target_status, target_stat.start_ticks
            )
            if target_check.name != process_stat.name:
                fail("App process identity changed before it could be signaled")
            requested_signal = signal.SIGTERM
        else:
            target_stat = process_stat
            requested_signal = getattr(signal, f"SIG{args.signal_name}")

        signal.pidfd_send_signal(target_fd, requested_signal)
        if args.signal_name == "APP_TERM":
            target_poll = select.poll()
            target_poll.register(target_fd, select.POLLIN)
            if target_poll.poll(APP_TERM_GRACE_MS):
                app_result = "graceful"
            else:
                try:
                    signal.pidfd_send_signal(target_fd, signal.SIGKILL)
                    app_result = "escalated"
                except ProcessLookupError:
                    # The exact pidfd target exited between the grace-period
                    # timeout and escalation, so no force signal was needed.
                    app_result = "graceful"
    except ProcessLookupError:
        target_pid = target_stat.pid if target_stat else args.pid
        fail(f"Process {target_pid} no longer exists")
    except PermissionError:
        target_pid = target_stat.pid if target_stat else args.pid
        fail(f"Permission denied while signaling process {target_pid}", 77)
    except OSError as error:
        target_pid = target_stat.pid if target_stat else args.pid
        fail(f"Could not signal process {target_pid}: {error}")
    finally:
        for opened_fd in opened_fds:
            os.close(opened_fd)

    result = f"signaled\t{target_stat.pid}\t{args.signal_name}"
    if app_result:
        result += f"\t{app_result}"
    print(result)


if __name__ == "__main__":
    main()
