#!/usr/bin/env python3
import os
import re
import signal
import subprocess
import sys
import time
import uuid

UUID = re.compile(r"^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$")
SESSION_ID = re.compile(r"^\$\d+$")
PANE_ID = re.compile(r"^%\d+$")
REAP_CGROUP = re.compile(
    r"^/user\.slice/user-\d+\.slice/user@\d+\.service/"
    r"(?:unsafe\.slice/unsafe-cld-\d+-\d+|agent\.slice/confine-agent-\d+-\d+|"
    r"human\.slice/human-agent-[0-9a-f-]{36})\.scope$"
)
TMUX_SCOPE = re.compile(
    r"^tmux-spawn-([0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12})\.scope$"
)
TMUX_DESCRIPTION = re.compile(r"^tmux child pane (\d+) launched by process (\d+)$")
TMUX_CGROUP = re.compile(
    rf"^/user\.slice/user-{os.getuid()}\.slice/user@{os.getuid()}\.service/"
    r"(?:app|human|agent)\.slice/tmux-spawn-[0-9a-f-]{36}\.scope$"
)
WATCHDOG_CGROUP = re.compile(rf"(?:{REAP_CGROUP.pattern}|{TMUX_CGROUP.pattern})")
INVOCATION_ID = re.compile(r"^[0-9a-f]{32}$")


def process_identity(pid):
    try:
        raw = open(f"/proc/{pid}/stat", encoding="utf-8").read()
        cgroup = open(f"/proc/{pid}/cgroup", encoding="utf-8").read().splitlines()
    except OSError:
        return None
    fields = raw[raw.rfind(")") + 2 :].split()
    unified = [line[3:] for line in cgroup if line.startswith("0::")]
    if len(unified) != 1:
        return None
    return int(fields[19]), unified[0]


def cgroup_value(cgroup_fd, name, key):
    fd = os.open(name, os.O_RDONLY | os.O_CLOEXEC, dir_fd=cgroup_fd)
    try:
        data = os.read(fd, 4096).decode()
    finally:
        os.close(fd)
    values = dict(line.split(None, 1) for line in data.splitlines())
    return values.get(key)


def set_frozen(cgroup_fd, frozen):
    fd = os.open("cgroup.freeze", os.O_WRONLY | os.O_CLOEXEC, dir_fd=cgroup_fd)
    try:
        os.write(fd, b"1" if frozen else b"0")
    finally:
        os.close(fd)


def empty_cgroup(cgroup_fd):
    fd = os.open("cgroup.kill", os.O_WRONLY | os.O_CLOEXEC, dir_fd=cgroup_fd)
    try:
        os.write(fd, b"1")
    finally:
        os.close(fd)


def watchdog_main(argv):
    if len(argv) not in (12, 15):
        return 64
    _, _, cgroup, parent_pid, parent_start, ready, disarm, mode, *recovery = argv
    runtime_dir = f"/run/user/{os.getuid()}/agent-session-reap"
    if not (
        WATCHDOG_CGROUP.fullmatch(cgroup)
        and parent_pid.isdigit()
        and parent_start.isdigit()
        and os.path.dirname(ready) == runtime_dir
        and os.path.dirname(disarm) == runtime_dir
        and ((mode == "session" and len(recovery) == 7) or (mode == "orphan" and len(recovery) == 4))
    ):
        return 64
    try:
        cgroup_fd = os.open(
            "/sys/fs/cgroup" + cgroup,
            os.O_RDONLY | os.O_DIRECTORY | os.O_CLOEXEC | os.O_NOFOLLOW,
        )
    except OSError:
        return 1
    try:
        self_identity = process_identity(os.getpid())
        if self_identity is None or self_identity[1] == cgroup or self_identity[1].startswith(cgroup + "/"):
            return 1
        try:
            ready_fd = os.open(
                ready,
                os.O_WRONLY | os.O_CREAT | os.O_TRUNC | os.O_CLOEXEC | os.O_NOFOLLOW,
                0o600,
            )
            os.write(ready_fd, b"READY\n")
            os.close(ready_fd)
        except OSError:
            return 1
        expected = int(parent_start)
        parent_pid_int = int(parent_pid)
        while True:
            identity = process_identity(parent_pid_int)
            if identity is None or identity[0] != expected:
                break
            time.sleep(0.05)

        try:
            disarm_fd = os.open(disarm, os.O_RDONLY | os.O_CLOEXEC | os.O_NOFOLLOW)
            try:
                disarmed = os.read(disarm_fd, 32) == b"DISARMED\n"
            finally:
                os.close(disarm_fd)
        except OSError:
            disarmed = False
        if disarmed:
            return 0

        try:
            set_frozen(cgroup_fd, True)
        except OSError:
            return 1
        deadline = time.monotonic() + 2
        while time.monotonic() < deadline:
            try:
                if cgroup_value(cgroup_fd, "cgroup.events", "frozen") == "1":
                    break
            except OSError:
                return 1
            time.sleep(0.01)
        else:
            return 1
        if not same_open_directory(cgroup_fd, "/sys/fs/cgroup" + cgroup):
            return 1

        recover = False
        if mode == "session":
            recover = tmux_session_is_absent(*recovery) is True
        else:
            unit, invocation, pane_pid, owner_pid = recovery
            state = tmux_scope_state(unit, systemctl="/usr/bin/systemctl")
            recover = (
                state == (invocation, cgroup, int(pane_pid), int(owner_pid))
                and tmux_pane_is_absent(int(owner_pid), int(pane_pid), tmux="/usr/bin/tmux") is True
            )
        if recover:
            try:
                empty_cgroup(cgroup_fd)
            except OSError:
                return 1
        return 0
    finally:
        try:
            set_frozen(cgroup_fd, False)
        except OSError:
            pass
        try:
            os.unlink(ready)
        except OSError:
            pass
        try:
            os.unlink(disarm)
        except OSError:
            pass
        os.close(cgroup_fd)


def disarm_watchdog(path):
    try:
        fd = os.open(path, os.O_WRONLY | os.O_TRUNC | os.O_CLOEXEC | os.O_NOFOLLOW)
        try:
            os.write(fd, b"DISARMED\n")
        finally:
            os.close(fd)
    except OSError:
        pass


def start_thaw_watchdog(cgroup, parent_start, recovery):
    runtime_dir = f"/run/user/{os.getuid()}/agent-session-reap"
    os.makedirs(runtime_dir, mode=0o700, exist_ok=True)
    os.chmod(runtime_dir, 0o700)
    nonce = uuid.uuid4().hex
    ready = f"{runtime_dir}/{nonce}.ready"
    disarm = f"{runtime_dir}/{nonce}.state"
    try:
        state_fd = os.open(
            disarm,
            os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_CLOEXEC | os.O_NOFOLLOW,
            0o600,
        )
        try:
            os.write(state_fd, b"ARMED\n")
        finally:
            os.close(state_fd)
    except OSError:
        return None
    result = subprocess.run(
        [
            "systemd-run",
            "--user",
            "--quiet",
            "--collect",
            "--service-type=exec",
            "--expand-environment=no",
            f"--unit=agent-reap-thaw-{nonce}",
            "--slice=app.slice",
            "--property=Restart=on-failure",
            "--property=RestartSec=100ms",
            "--",
            sys.executable,
            os.path.realpath(__file__),
            "--watchdog",
            cgroup,
            str(os.getpid()),
            str(parent_start),
            ready,
            disarm,
            *recovery,
        ],
        check=False,
        capture_output=True,
        text=True,
        timeout=5,
    )
    if result.returncode != 0 or result.stderr:
        disarm_watchdog(disarm)
        return None
    deadline = time.monotonic() + 2
    while time.monotonic() < deadline:
        try:
            if open(ready, encoding="utf-8").read() == "READY\n":
                os.unlink(ready)
                return disarm
        except OSError:
            pass
        time.sleep(0.01)
    disarm_watchdog(disarm)
    return None


def tmux_scope_state(unit, systemctl="systemctl"):
    if not TMUX_SCOPE.fullmatch(unit):
        return None
    result = subprocess.run(
        [
            systemctl,
            "--user",
            "show",
            unit,
            "-p",
            "Id",
            "-p",
            "Description",
            "-p",
            "ActiveState",
            "-p",
            "InvocationID",
            "-p",
            "ControlGroup",
        ],
        check=False,
        capture_output=True,
        text=True,
        timeout=5,
    )
    if result.returncode != 0 or result.stderr:
        return None
    values = dict(line.split("=", 1) for line in result.stdout.splitlines() if "=" in line)
    description = TMUX_DESCRIPTION.fullmatch(values.get("Description", ""))
    cgroup = values.get("ControlGroup", "")
    if not (
        values.get("Id") == unit
        and values.get("ActiveState") == "active"
        and INVOCATION_ID.fullmatch(values.get("InvocationID", ""))
        and description
        and TMUX_CGROUP.fullmatch(cgroup)
        and cgroup.endswith("/" + unit)
    ):
        return None
    pane_pid, owner_pid = (int(value) for value in description.groups())
    if pane_pid <= 0 or owner_pid <= 0:
        return None
    return values.get("InvocationID"), cgroup, pane_pid, owner_pid


def pid_is_absent(pid):
    try:
        pidfd = os.pidfd_open(pid)
    except ProcessLookupError:
        return True
    except (AttributeError, OSError):
        return None
    os.close(pidfd)
    return False


def tmux_socket_paths(pid):
    try:
        comm = open(f"/proc/{pid}/comm", encoding="utf-8").read().strip()
        fd_names = os.listdir(f"/proc/{pid}/fd")
        unix_lines = open("/proc/net/unix", encoding="utf-8").read().splitlines()[1:]
    except OSError:
        return None
    if comm != "tmux: server":
        return None
    inodes = set()
    for name in fd_names:
        try:
            target = os.readlink(f"/proc/{pid}/fd/{name}")
        except OSError:
            return None
        match = re.fullmatch(r"socket:\[(\d+)\]", target)
        if match:
            inodes.add(match.group(1))
    paths = set()
    for line in unix_lines:
        fields = line.split()
        if len(fields) >= 8 and fields[6] in inodes and fields[7].startswith("/"):
            paths.add(fields[7])
    return sorted(paths) or None


def tmux_pane_is_absent(owner_pid, pane_pid, tmux="tmux"):
    owner_absent = pid_is_absent(owner_pid)
    if owner_absent is True:
        return True
    if owner_absent is not False:
        return None
    sockets = tmux_socket_paths(owner_pid)
    if sockets is None:
        return None
    saw_server = False
    for socket in sockets:
        result = subprocess.run(
            [tmux, "-S", socket, "list-panes", "-a", "-F", "#{pid}|#{pane_pid}"],
            check=False,
            capture_output=True,
            text=True,
            timeout=5,
        )
        if result.returncode != 0 or result.stderr:
            return None
        for line in result.stdout.splitlines():
            fields = line.split("|")
            if len(fields) != 2 or not all(value.isdigit() for value in fields):
                return None
            if int(fields[0]) != owner_pid:
                return None
            saw_server = True
            if int(fields[1]) == pane_pid:
                return False
    return True if saw_server else None


def tmux_session_is_absent(
    socket, server_pid, session_id, created, token, pane_id, pane_pid, tmux="/usr/bin/tmux"
):
    actual = (
        "#{pid}|#{session_id}|#{session_created}|#{@agent_reap_identity}|"
        "#{pane_id}|#{pane_pid}"
    )
    result = subprocess.run(
        [tmux, "-S", socket, "list-panes", "-a", "-F", actual],
        check=False,
        capture_output=True,
        text=True,
        timeout=5,
    )
    if result.returncode == 0 and not result.stderr:
        lines = result.stdout.splitlines()
        if not lines:
            return None
        for line in lines:
            fields = line.split("|")
            if not (
                len(fields) == 6
                and fields[0].isdigit()
                and fields[0] == server_pid
                and SESSION_ID.fullmatch(fields[1])
                and fields[2].isdigit()
                and (fields[3] == "" or UUID.fullmatch(fields[3]))
                and PANE_ID.fullmatch(fields[4])
                and fields[5].isdigit()
            ):
                return None
            if fields[1] == session_id:
                return False
        return True
    return True if pid_is_absent(int(server_pid)) is True else None


def same_open_directory(fd, path):
    try:
        opened = os.fstat(fd)
        current = os.stat(path, follow_symlinks=False)
    except OSError:
        return False
    return opened.st_dev == current.st_dev and opened.st_ino == current.st_ino


def reap_orphan_scope(unit, dry_run=False):
    state = tmux_scope_state(unit)
    if state is None:
        print(f"agent-session-reap: kept {unit}: invalid scope identity", file=sys.stderr)
        return False
    invocation, cgroup, _, owner_pid = state
    if tmux_pane_is_absent(owner_pid, state[2]) is not True:
        return False
    if dry_run:
        print(f"would reap orphan scope {unit}")
        return True

    cgroup_path = "/sys/fs/cgroup" + cgroup
    try:
        cgroup_fd = os.open(
            cgroup_path,
            os.O_RDONLY | os.O_DIRECTORY | os.O_CLOEXEC | os.O_NOFOLLOW,
        )
    except OSError as error:
        print(f"agent-session-reap: kept {unit}: cannot bind scope: {error}", file=sys.stderr)
        return False

    if not same_open_directory(cgroup_fd, cgroup_path):
        os.close(cgroup_fd)
        print(f"agent-session-reap: kept {unit}: scope directory changed", file=sys.stderr)
        return False
    self_identity = process_identity(os.getpid())
    if self_identity is None or self_identity[1] == cgroup or self_identity[1].startswith(cgroup + "/"):
        os.close(cgroup_fd)
        print(f"agent-session-reap: kept {unit}: helper belongs to target scope", file=sys.stderr)
        return False

    parent_start = process_identity(os.getpid())[0]
    recovery = ["orphan", unit, invocation, str(state[2]), str(owner_pid)]
    watchdog = start_thaw_watchdog(cgroup, parent_start, recovery)
    if watchdog is None:
        os.close(cgroup_fd)
        print(f"agent-session-reap: kept {unit}: thaw watchdog did not become ready", file=sys.stderr)
        return False

    try:
        try:
            set_frozen(cgroup_fd, True)
        except OSError as error:
            print(f"agent-session-reap: kept {unit}: cannot freeze scope: {error}", file=sys.stderr)
            return False
        deadline = time.monotonic() + 2
        while time.monotonic() < deadline:
            try:
                if cgroup_value(cgroup_fd, "cgroup.events", "frozen") == "1":
                    break
            except OSError as error:
                print(f"agent-session-reap: kept {unit}: cannot read scope state: {error}", file=sys.stderr)
                return False
            time.sleep(0.01)
        else:
            print(f"agent-session-reap: kept {unit}: scope did not freeze", file=sys.stderr)
            return False

        if not same_open_directory(cgroup_fd, cgroup_path):
            print(f"agent-session-reap: kept {unit}: scope directory changed while freezing", file=sys.stderr)
            return False
        if tmux_scope_state(unit) != (invocation, cgroup, state[2], owner_pid):
            print(f"agent-session-reap: kept {unit}: scope identity changed while freezing", file=sys.stderr)
            return False
        if tmux_pane_is_absent(owner_pid, state[2]) is not True:
            print(f"agent-session-reap: kept {unit}: tmux pane appeared while freezing", file=sys.stderr)
            return False

        try:
            empty_cgroup(cgroup_fd)
        except OSError as error:
            print(f"agent-session-reap: kept {unit}: cannot empty scope: {error}", file=sys.stderr)
            return False
        print(f"reaped orphan scope {unit}")
        return True
    finally:
        disarm_watchdog(watchdog)
        try:
            set_frozen(cgroup_fd, False)
        except OSError:
            pass
        os.close(cgroup_fd)


def orphan_scopes_main(dry_run=False):
    result = subprocess.run(
        [
            "systemctl",
            "--user",
            "list-units",
            "--all",
            "--type=scope",
            "--plain",
            "--no-legend",
            "tmux-spawn-*.scope",
        ],
        check=False,
        capture_output=True,
        text=True,
        timeout=10,
    )
    if result.returncode != 0 or result.stderr:
        print("agent-session-reap: cannot enumerate tmux scopes", file=sys.stderr)
        return 1
    reaped = 0
    invalid = False
    for line in result.stdout.splitlines():
        fields = line.split()
        if not fields:
            continue
        unit = fields[0]
        if not TMUX_SCOPE.fullmatch(unit):
            print(f"agent-session-reap: ignored malformed tmux scope name: {unit}", file=sys.stderr)
            invalid = True
            continue
        state = tmux_scope_state(unit)
        if state is None:
            print(f"agent-session-reap: kept {unit}: invalid scope identity", file=sys.stderr)
            invalid = True
            continue
        if reap_orphan_scope(unit, dry_run=dry_run):
            reaped += 1
    if reaped:
        action = "would reap" if dry_run else "reaped"
        print(f"{action} {reaped} orphan tmux scope(s)")
    return 1 if invalid else 0


def kept(reason):
    print("KEPT")
    print(f"agent-session-reap: {reason}", file=sys.stderr)
    return 0


def main(argv):
    if len(argv) != 11:
        return kept("invalid helper arguments")
    socket, server_pid, session_id, created, token, pane_id, pane_pid, pane_start, runtime_pid, runtime_start = argv[1:]
    if not (
        server_pid.isdigit()
        and created.isdigit()
        and pane_pid.isdigit()
        and pane_start.isdigit()
        and runtime_pid.isdigit()
        and runtime_start.isdigit()
        and SESSION_ID.fullmatch(session_id)
        and PANE_ID.fullmatch(pane_id)
        and UUID.fullmatch(token)
    ):
        return kept("invalid tmux identity")

    pane_pid_int = int(pane_pid)
    pane_start_int = int(pane_start)
    try:
        pidfd = os.pidfd_open(pane_pid_int)
    except (AttributeError, OSError) as error:
        return kept(f"cannot open pane pidfd: {error}")

    identity = process_identity(pane_pid_int)
    if identity is None or identity[0] != pane_start_int:
        os.close(pidfd)
        return kept("pane process identity changed")
    cgroup = identity[1]
    if not REAP_CGROUP.fullmatch(cgroup):
        os.close(pidfd)
        return kept("pane does not own an eligible session scope")

    self_identity = process_identity(os.getpid())
    if self_identity is None or self_identity[1] == cgroup or self_identity[1].startswith(cgroup + "/"):
        os.close(pidfd)
        return kept("helper belongs to the target scope")

    try:
        cgroup_fd = os.open(
            "/sys/fs/cgroup" + cgroup,
            os.O_RDONLY | os.O_DIRECTORY | os.O_CLOEXEC | os.O_NOFOLLOW,
        )
        procs_fd = os.open("cgroup.procs", os.O_RDONLY | os.O_CLOEXEC, dir_fd=cgroup_fd)
        try:
            root_pids = os.read(procs_fd, 1_048_576).decode().split()
        finally:
            os.close(procs_fd)
        if str(pane_pid_int) not in root_pids:
            raise OSError("pane leader is not in scope root")
    except OSError as error:
        os.close(pidfd)
        try:
            os.close(cgroup_fd)
        except (NameError, OSError):
            pass
        return kept(f"cannot bind pane scope: {error}")

    runtime_pid_int = int(runtime_pid)
    runtime_start_int = int(runtime_start)
    try:
        runtime_pidfd = os.pidfd_open(runtime_pid_int)
        runtime_identity = process_identity(runtime_pid_int)
        if runtime_identity is None or runtime_identity[0] != runtime_start_int:
            raise OSError("runtime process identity changed")
        runtime_cgroup = runtime_identity[1]
        if runtime_cgroup != cgroup and not runtime_cgroup.startswith(cgroup + "/"):
            raise OSError("runtime process left pane scope")
    except (AttributeError, OSError) as error:
        try:
            os.close(runtime_pidfd)
        except (NameError, OSError):
            pass
        os.close(cgroup_fd)
        os.close(pidfd)
        return kept(f"cannot bind runtime process: {error}")

    parent_start = process_identity(os.getpid())[0]
    recovery = ["session", socket, server_pid, session_id, created, token, pane_id, pane_pid]
    watchdog = start_thaw_watchdog(cgroup, parent_start, recovery)
    if watchdog is None:
        os.close(runtime_pidfd)
        os.close(cgroup_fd)
        os.close(pidfd)
        return kept("independent thaw watchdog did not become ready")

    kill_fd = None
    recovery_needed = False
    try:
        try:
            set_frozen(cgroup_fd, True)
        except OSError as error:
            return kept(f"cannot freeze pane scope: {error}")

        deadline = time.monotonic() + 2
        while time.monotonic() < deadline:
            try:
                if cgroup_value(cgroup_fd, "cgroup.events", "frozen") == "1":
                    break
            except OSError as error:
                return kept(f"cannot read pane scope state: {error}")
            time.sleep(0.01)
        else:
            return kept("pane scope did not freeze")

        try:
            signal.pidfd_send_signal(pidfd, 0)
        except (AttributeError, OSError, ProcessLookupError) as error:
            return kept(f"pane exited before atomic close: {error}")
        identity = process_identity(pane_pid_int)
        if identity != (pane_start_int, cgroup):
            return kept("pane process identity changed while freezing")
        try:
            signal.pidfd_send_signal(runtime_pidfd, 0)
        except (AttributeError, OSError, ProcessLookupError) as error:
            return kept(f"runtime exited before atomic close: {error}")
        runtime_identity = process_identity(runtime_pid_int)
        if runtime_identity is None or runtime_identity[0] != runtime_start_int:
            return kept("runtime process identity changed while freezing")
        if runtime_identity[1] != runtime_cgroup:
            return kept("runtime process cgroup changed while freezing")
        try:
            kill_fd = os.open("cgroup.kill", os.O_WRONLY | os.O_CLOEXEC, dir_fd=cgroup_fd)
        except OSError as error:
            return kept(f"cannot bind scope cleanup: {error}")

        actual = (
            "#{pid}|#{session_id}|#{session_created}|#{@agent_reap_identity}|"
            "#{pane_id}|#{pane_pid}|#{session_attached}|#{window_linked}"
        )
        expected = f"{server_pid}|{session_id}|{created}|{token}|{pane_id}|{pane_pid}|0|0"
        recovery_needed = True
        result = subprocess.run(
            [
                "tmux",
                "-S",
                socket,
                "if-shell",
                "-t",
                pane_id,
                "-F",
                f"#{{==:{actual},{expected}}}",
                f"kill-session -t {session_id}",
                "display-message -p KEPT",
            ],
            check=False,
            capture_output=True,
            text=True,
            timeout=5,
        )
        if result.returncode != 0 or result.stderr or result.stdout not in ("", "KEPT\n"):
            return kept("tmux atomic close failed")
        if result.stdout == "KEPT\n":
            recovery_needed = False
            return kept("attachment or tmux identity changed")
        try:
            while True:
                try:
                    os.write(kill_fd, b"1")
                    break
                except InterruptedError:
                    continue
        except OSError as error:
            print(f"agent-session-reap: tmux closed; watchdog will retry scope cleanup: {error}", file=sys.stderr)
            print("REAPED")
            return 0
        recovery_needed = False
        print("REAPED")
        return 0
    finally:
        if not recovery_needed:
            disarm_watchdog(watchdog)
        if kill_fd is not None:
            os.close(kill_fd)
        try:
            set_frozen(cgroup_fd, False)
        except OSError:
            pass
        os.close(cgroup_fd)
        os.close(runtime_pidfd)
        os.close(pidfd)


if __name__ == "__main__":
    if len(sys.argv) > 1 and sys.argv[1] == "--watchdog":
        raise SystemExit(watchdog_main(sys.argv))
    if len(sys.argv) in (2, 3) and sys.argv[1] == "--orphan-scopes":
        if len(sys.argv) == 3 and sys.argv[2] != "--dry-run":
            raise SystemExit(64)
        raise SystemExit(orphan_scopes_main(dry_run=len(sys.argv) == 3))
    raise SystemExit(main(sys.argv))
