#!/usr/bin/python3
"""mood-kill — force stop apps (runs as the user, never root).

  mood-kill <app-id> [--pid N]...   kill every process of an app, print JSON {"killed": n}
  mood-kill --pid N                 kill the app that owns process N (window of a frozen app)
  mood-kill --list                  JSON list of running apps [{id, name, pids, mem_mb}]

How an app's processes are found (all of them are tried):
  1. its systemd scope: mood-launch starts every app as app-mood-<id>-<n>.scope (Description "Mood app <id>")
  2. Flatpak: flatpak kill <id> (flatpak moves apps into its own app-flatpak-*.scope)
  3. the cgroup of a window's PID, when it is an app-*.scope
  4. leftovers started before 1.4: mood-app/mood-browser command lines, or the .desktop executable's name
Session processes (Wayfire, the Mood shell, systemd, dbus, pipewire…) are never touched."""
import json
import os
import re
import shutil
import signal
import subprocess
import sys
import time
from pathlib import Path

UID = os.getuid()
ME = os.getpid()
PROTECT = {"wayfire", "labwc", "systemd", "dbus-daemon", "dbus-broker", "pipewire", "wireplumber", "pipewire-pulse",
           "greetd", "mako", "swayidle", "gtklock", "Xwayland", "mood-shell", "mood-session", "sd-pam", "gvfsd",
           "xdg-desktop-portal", "xdg-document-portal", "xdg-permission-store", "at-spi-bus-launcher", "at-spi2-registryd"}
GENERIC_EXE = {"python3", "python", "sh", "bash", "dash", "env", "flatpak", "bwrap", "java", "node", "electron", "wine",
               "mono", "perl", "ruby", "kitty", "xdg-open", "gio"}


def sh(*cmd, timeout=8):
    try:
        r = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout)
        return r.returncode, r.stdout
    except (OSError, subprocess.TimeoutExpired):
        return 1, ""


def safe_id(app_id):
    return re.sub(r"[^A-Za-z0-9_.]", "_", app_id)[:80] or "app"


# ------------------------------------------------------------------ process table
def procs():
    out = {}
    for d in Path("/proc").iterdir():
        if not d.name.isdigit():
            continue
        try:
            if d.stat().st_uid != UID:
                continue
            stat = (d / "stat").read_text()
            ppid = int(stat[stat.rfind(")") + 2:].split()[1])
            cmd = (d / "cmdline").read_bytes().split(b"\0")
            out[int(d.name)] = {
                "ppid": ppid, "comm": (d / "comm").read_text().strip(),
                "argv": [c.decode(errors="replace") for c in cmd if c],
                "cgroup": (d / "cgroup").read_text().strip().rsplit("/", 1)[-1],
            }
        except (OSError, ValueError, IndexError):
            continue
    return out


def protected(p):
    if p["comm"] in PROTECT:
        return True
    a = " ".join(p["argv"][:3])
    return "mood-shell" in a or "mood-session" in a or "moodtop-wm" in a or "mood-bt-agent" in a


def tree(pids, table):
    """pids + all their descendants (only user-owned, unprotected, never ourselves)."""
    kids = {}
    for pid, p in table.items():
        kids.setdefault(p["ppid"], []).append(pid)
    out, todo = set(), list(pids)
    while todo:
        pid = todo.pop()
        if pid in out or pid == ME or pid not in table or protected(table[pid]):
            continue
        out.add(pid)
        todo.extend(kids.get(pid, []))
    return out


# ------------------------------------------------------------------ scopes
def scopes():
    """{unit: {"id": app id, "desc": description, "mem": bytes}} for app-*.scope units in the user manager."""
    code, out = sh("systemctl", "--user", "show", "--type=scope", "-p", "Id,Description,MemoryCurrent,ActiveState", "app-*.scope")
    res = {}
    for block in out.strip().split("\n\n") if code == 0 else []:
        kv = dict(l.split("=", 1) for l in block.splitlines() if "=" in l)
        unit = kv.get("Id", "")
        if not unit.startswith("app-") or kv.get("ActiveState") not in ("active", "activating"):
            continue
        desc = kv.get("Description", "")
        m = re.match(r"Mood app (.+)$", desc)
        if m:
            aid = m.group(1)
        else:
            m = re.match(r"app-flatpak-(.+)-\d+\.scope$", unit)
            aid = m.group(1) if m else ""
        mem = kv.get("MemoryCurrent", "")
        res[unit] = {"id": aid, "desc": desc, "mem": int(mem) if mem.isdigit() else 0}
    return res


def kill_units(units):
    if not units:
        return
    sh("systemctl", "--user", "kill", "--signal=SIGKILL", *units)
    sh("systemctl", "--user", "stop", "--no-block", *units)


# ------------------------------------------------------------------ desktop executable
def desktop_exe(app_id):
    try:
        import gi
        gi.require_version("Gio", "2.0")
        from gi.repository import Gio
        info = Gio.DesktopAppInfo.new(app_id if app_id.endswith(".desktop") else app_id + ".desktop")
        exe = info.get_executable() if info else None
    except Exception:
        exe = None
    if not exe:
        return None
    path = shutil.which(exe) or exe
    try:
        path = os.path.realpath(path)
    except OSError:
        pass
    name = os.path.basename(path)
    return None if name in GENERIC_EXE else name


def legacy_pids(app_id, table):
    """Processes of an app that was started without a scope (before 1.4, or from a terminal)."""
    out = set()
    mood_app = Path(f"/usr/lib/mood/apps/{app_id}/manifest.json").exists()
    exe = None if mood_app or app_id in ("browser", "terminal", "zapos") else desktop_exe(app_id)
    for pid, p in table.items():
        a = p["argv"]
        if len(a) >= 3 and a[1].endswith("/mood-app") and a[2] == app_id:
            out.add(pid)
        elif len(a) >= 2 and a[0].endswith("mood-app") and a[1] == app_id:
            out.add(pid)
        elif app_id == "browser" and any(x.endswith("mood-browser.py") for x in a[:2]):
            out.add(pid)
        elif exe and (p["comm"] == exe[:15] or (a and os.path.basename(a[0]) == exe)):
            out.add(pid)
    return out


# ------------------------------------------------------------------ actions
def force_stop(app_id=None, pids=()):
    table = procs()
    units, targets = set(), set()
    sc = scopes()
    if app_id:
        units |= {u for u, s in sc.items() if s["id"] == app_id or s["id"] == safe_id(app_id)}
        if "." in app_id and shutil.which("flatpak"):
            code, _ = sh("flatpak", "info", app_id)
            if code == 0:
                sh("flatpak", "kill", app_id)
                units |= {u for u in sc if u.startswith(f"app-flatpak-{app_id}-")}
        targets |= legacy_pids(app_id, table)
    for pid in pids:
        p = table.get(pid)
        if not p:
            continue
        if p["cgroup"].startswith("app-") and p["cgroup"].endswith(".scope"):
            units.add(p["cgroup"])
        targets.add(pid)
    # everything inside the scopes we're about to kill is counted too
    for pid, p in table.items():
        if p["cgroup"] in units and not protected(p):
            targets.add(pid)
    targets = tree(targets, table)
    kill_units(sorted(units))
    for pid in targets:
        try:
            os.kill(pid, signal.SIGKILL)
        except OSError:
            pass
    # give the kernel a moment so a relaunch right after doesn't race the dying single-instance name
    deadline = time.time() + 1.5
    while time.time() < deadline and any(Path(f"/proc/{p}").exists() for p in targets):
        time.sleep(0.05)
    return {"killed": len(targets), "units": sorted(units)}


def running():
    table = procs()
    sc = scopes()
    apps = {}
    for unit, s in sc.items():
        if not s["id"]:
            continue
        pids = [pid for pid, p in table.items() if p["cgroup"] == unit]
        if not pids:
            continue
        a = apps.setdefault(s["id"], {"id": s["id"], "pids": [], "mem_mb": 0, "units": []})
        a["pids"] += pids
        a["units"].append(unit)
        a["mem_mb"] += s["mem"] // (1024 * 1024)
    return sorted(apps.values(), key=lambda a: -a["mem_mb"])


def main(argv):
    if argv[:1] == ["--list"]:
        print(json.dumps(running()))
        return 0
    app_id, pids, i = None, [], 0
    while i < len(argv):
        if argv[i] == "--pid" and i + 1 < len(argv):
            if argv[i + 1].isdigit():
                pids.append(int(argv[i + 1]))
            i += 2
            continue
        if not argv[i].startswith("-") and app_id is None:
            app_id = argv[i]
        i += 1
    if not app_id and not pids:
        print(__doc__, file=sys.stderr)
        return 2
    print(json.dumps(force_stop(app_id, pids)))
    return 0


if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
