from __future__ import annotations

import os
import signal
import time
from pathlib import Path


def pick_fastest_grower(before: dict[str, int], after: dict[str, int], protect: set[str]) -> dict | None:
    best = None
    for name, kb in after.items():
        if name in protect:
            continue
        growth = kb - before.get(name, 0)
        if growth <= 0:
            continue
        if best is None or growth > best["growth_kb"]:
            best = {"name": name, "growth_kb": growth, "rss_kb": kb}
    return best


def _proc_name(pid_dir: Path) -> str:
    try:
        return os.path.basename(os.readlink(pid_dir / "exe"))
    except OSError:
        try:
            return (pid_dir / "comm").read_text().strip()
        except OSError:
            return pid_dir.name


def _start_time(stat: str) -> str:
    end = stat.rfind(")")
    parts = stat[end + 2:].split()
    return parts[19] if len(parts) > 19 else ""


def scan_processes(proc_root: Path = Path("/proc")) -> dict[int, dict[str, str]]:
    out = {}
    for entry in proc_root.iterdir():
        if not entry.name.isdigit():
            continue
        try:
            out[int(entry.name)] = {
                "exe": _proc_name(entry),
                "start_time": _start_time((entry / "stat").read_text()),
                "cgroup": (entry / "cgroup").read_text().strip(),
            }
        except OSError:
            continue
    return out


def rss_by_group(proc_root: Path = Path("/proc")) -> tuple[dict[str, int], dict[str, list[dict[str, str | int]]]]:
    totals: dict[str, int] = {}
    pids: dict[str, list[dict[str, str | int]]] = {}
    for entry in proc_root.iterdir():
        if not entry.name.isdigit():
            continue
        try:
            name = _proc_name(entry)
            status = (entry / "status").read_text()
            rss = 0
            for line in status.splitlines():
                if line.startswith("VmRSS:"):
                    rss = int(line.split()[1])
                    break
            stat = (entry / "stat").read_text()
            cgroup = (entry / "cgroup").read_text().strip()
        except OSError:
            continue
        totals[name] = totals.get(name, 0) + rss
        pids.setdefault(name, []).append({"pid": int(entry.name), "start_time": _start_time(stat), "cgroup": cgroup})
    return totals, pids


def rescan_culprit(protect: set[str], proc_root: Path = Path("/proc"), delay: float = 1.0) -> dict | None:
    before, _ = rss_by_group(proc_root)
    time.sleep(delay)
    after, pids = rss_by_group(proc_root)
    culprit = pick_fastest_grower(before, after, protect)
    if culprit:
        culprit["pids"] = pids.get(culprit["name"], [])
    return culprit


def safe_signal_plan(target: dict, live: dict[int, dict[str, str]], protect: set[str]) -> list[int]:
    if target.get("name") in protect:
        return []
    allowed = []
    for item in target.get("pids", []):
        pid = int(item.get("pid", -1))
        current = live.get(pid)
        if not current or current.get("exe") in protect:
            continue
        if current.get("start_time") != str(item.get("start_time")):
            continue
        if current.get("cgroup") != item.get("cgroup"):
            continue
        allowed.append(pid)
    return allowed


def terminate_target(target: dict, protect: set[str]) -> list[int]:
    first = safe_signal_plan(target, scan_processes(), protect)
    for pid in first:
        try:
            os.kill(pid, signal.SIGTERM)
        except ProcessLookupError:
            pass
    time.sleep(3)
    second = safe_signal_plan(target, scan_processes(), protect)
    for pid in second:
        try:
            os.kill(pid, signal.SIGKILL)
        except ProcessLookupError:
            pass
    return second
