#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Sentinel 探针 —— 纯标准库、单文件、兼容 Python >= 3.6。

采集 CPU/内存/磁盘/网络/进程 等指标，定时 POST 到 Sentinel 服务端；
上报间隔、进程清单、最新版本号全部由服务端在响应里下发（服务端集中配置）。
自动更新：响应中的版本号不同 → 下载/校验/自替换/退出，由 systemd 拉起新版。

红线（见服务端仓库 CLAUDE.md）：只用标准库；不用 subprocess 的
capture_output=/text=（3.7+）等新参数；配置在 /opt/sentinel-agent/config.json。
"""

import hashlib
import http.client
import json
import os
import platform
import py_compile
import re
import shutil
import socket
import sys
import time
import urllib.parse
import urllib.request

__VERSION__ = "1.1.0"

CONFIG_PATH = "/opt/sentinel-agent/config.json"
UPDATE_MARKER = "/opt/sentinel-agent/.last_update_attempt"
UPDATE_COOLDOWN = 600  # 秒；更新失败后的最短重试间隔，防更新风暴
SELF_PATH = os.path.abspath(__file__)

# 统计流量时排除的接口前缀：回环与虚拟设备（隧道流量已计入物理网卡，再计一遍会翻倍）
IFACE_EXCLUDE = ("lo", "veth", "docker", "br-", "virbr", "tun", "tap", "wg", "zt", "tailscale")
REAL_FS = {"ext2", "ext3", "ext4", "xfs", "btrfs", "zfs", "f2fs", "vfat", "exfat",
           "ntfs", "fuseblk", "jfs", "reiserfs"}
MOUNT_SKIP_PREFIX = ("/var/lib/docker", "/snap", "/boot/efi", "/run")


def log(msg):
    sys.stdout.write("%s %s\n" % (time.strftime("%Y-%m-%d %H:%M:%S"), msg))
    sys.stdout.flush()


def load_config():
    try:
        with open(CONFIG_PATH) as f:
            cfg = json.load(f)
    except Exception as exc:
        log("无法读取配置 %s: %s" % (CONFIG_PATH, exc))
        sys.exit(2)
    for k in ("server_url", "token"):
        if not cfg.get(k):
            log("配置缺少 %s" % k)
            sys.exit(2)
    cfg.setdefault("name", socket.gethostname())
    cfg.setdefault("interval", 10)
    return cfg


# ---------------------------------------------------------------- 采集
def read_proc_stat():
    with open("/proc/stat") as f:
        parts = f.readline().split()
    vals = [int(x) for x in parts[1:9]]  # user nice system idle iowait irq softirq steal
    idle = vals[3] + vals[4]
    return sum(vals), idle


def cpu_percent(prev, cur):
    if prev is None:
        return None
    dt, di = cur[0] - prev[0], cur[1] - prev[1]
    if dt <= 0:
        return None
    return round(100.0 * (1.0 - float(di) / dt), 1)


def read_meminfo():
    info = {}
    with open("/proc/meminfo") as f:
        for line in f:
            fields = line.split()
            if len(fields) >= 2:
                info[fields[0].rstrip(":")] = int(fields[1]) * 1024
    total = info.get("MemTotal", 0)
    avail = info.get("MemAvailable")
    if avail is None:  # 老内核（<3.14）没有 MemAvailable
        avail = info.get("MemFree", 0) + info.get("Buffers", 0) + info.get("Cached", 0)
    used = max(total - avail, 0)
    mem = {"total": total, "used": used,
           "percent": round(100.0 * used / total, 1) if total else 0.0}
    st, sf = info.get("SwapTotal", 0), info.get("SwapFree", 0)
    return mem, {"total": st, "used": max(st - sf, 0)}


def read_disks():
    out, seen = [], set()
    try:
        with open("/proc/mounts") as f:
            lines = f.read().splitlines()
    except OSError:
        return out
    for line in lines:
        parts = line.split()
        if len(parts) < 3:
            continue
        dev, mnt, fs = parts[0], parts[1], parts[2]
        if fs not in REAL_FS or dev in seen:
            continue
        mnt = mnt.replace("\\040", " ").replace("\\011", "\t")
        if mnt.startswith(MOUNT_SKIP_PREFIX):
            continue
        seen.add(dev)
        try:
            st = os.statvfs(mnt)
        except OSError:
            continue
        total = st.f_blocks * st.f_frsize
        if total <= 0:
            continue
        used = (st.f_blocks - st.f_bfree) * st.f_frsize
        avail = st.f_bavail * st.f_frsize
        pct = round(100.0 * used / (used + avail), 1) if used + avail > 0 else 0.0
        out.append({"mount": mnt, "total": total, "used": used, "percent": pct})
    out.sort(key=lambda d: d["mount"])
    return out[:8]


def read_net():
    rx = tx = 0
    try:
        with open("/proc/net/dev") as f:
            lines = f.read().splitlines()[2:]
    except OSError:
        return 0, 0
    for line in lines:
        if ":" not in line:
            continue
        name, rest = line.split(":", 1)
        if name.strip().startswith(IFACE_EXCLUDE):
            continue
        fields = rest.split()
        try:
            rx += int(fields[0])
            tx += int(fields[8])
        except (IndexError, ValueError):
            continue
    return rx, tx


def read_boot_id():
    try:
        with open("/proc/sys/kernel/random/boot_id") as f:
            return f.read().strip()
    except OSError:
        return "unknown"


def read_uptime():
    try:
        with open("/proc/uptime") as f:
            return int(float(f.read().split()[0]))
    except (OSError, ValueError):
        return 0


def read_os():
    try:
        with open("/etc/os-release") as f:
            for line in f:
                if line.startswith("PRETTY_NAME="):
                    return line.split("=", 1)[1].strip().strip('"')
    except OSError:
        pass
    return platform.system()


def check_processes(names):
    """按进程名（/proc/*/comm，注意 15 字符截断）或命令行子串判断存活。"""
    result = dict((n, False) for n in names)
    if not names:
        return result
    my_pid = str(os.getpid())
    for pid in os.listdir("/proc"):
        if not pid.isdigit() or pid == my_pid:
            continue
        try:
            with open("/proc/%s/comm" % pid) as f:
                comm = f.read().strip()
            with open("/proc/%s/cmdline" % pid, "rb") as f:
                cmdline = f.read().replace(b"\0", b" ").decode("utf-8", "replace")
        except (OSError, IOError):
            continue
        for n in names:
            if not result[n] and (n == comm or (n and n in cmdline)):
                result[n] = True
    return result


def collect(state, cfg):
    cur_stat = read_proc_stat()
    cpu = cpu_percent(state.get("prev_stat"), cur_stat)
    state["prev_stat"] = cur_stat
    rx, tx = read_net()
    mono = time.monotonic()
    rx_rate = tx_rate = None
    prev = state.get("prev_net")
    if prev is not None:
        dt = mono - prev[2]
        if dt > 0 and rx >= prev[0] and tx >= prev[1]:
            rx_rate = round((rx - prev[0]) / dt, 1)
            tx_rate = round((tx - prev[1]) / dt, 1)
    state["prev_net"] = (rx, tx, mono)
    mem, swap = read_meminfo()
    try:
        load1 = round(os.getloadavg()[0], 2)
    except OSError:
        load1 = None
    return {
        "name": cfg["name"], "version": __VERSION__, "boot_id": read_boot_id(),
        "os": read_os(), "kernel": platform.release(), "arch": platform.machine(),
        "uptime_s": read_uptime(), "load1": load1, "ts": time.time(),
        "cpu": {"percent": cpu, "cores": os.cpu_count() or 1},
        "mem": mem, "swap": swap, "disks": read_disks(),
        "net": {"rx_bytes": rx, "tx_bytes": tx, "rx_rate": rx_rate, "tx_rate": tx_rate},
        "processes": check_processes(state.get("processes") or []),
    }


# ---------------------------------------------------------------- 上报与更新
class ServerLink(object):
    """与服务端的持久 HTTP(S) 连接（keep-alive）。

    v1.1.0 起不再每次上报新建 TCP+TLS（每次握手约 5KB，是探针流量的大头）；
    连接被对端闲置关闭时自动重连一次。证书校验与 urllib 相同（系统 CA）。
    """

    def __init__(self, base_url, timeout=10):
        u = urllib.parse.urlparse(base_url)
        self.https = (u.scheme == "https")
        self.host = u.hostname
        self.port = u.port or (443 if self.https else 80)
        self.prefix = u.path.rstrip("/")
        self.timeout = timeout
        self.conn = None

    def close(self):
        if self.conn is not None:
            try:
                self.conn.close()
            except Exception:
                pass
            self.conn = None

    def post_json(self, path, obj, headers):
        body = json.dumps(obj).encode("utf-8")
        last_exc = None
        for attempt in (1, 2):  # 第 1 次失败多半是闲置连接被对端关了：重连再试一次
            try:
                if self.conn is None:
                    cls = http.client.HTTPSConnection if self.https else http.client.HTTPConnection
                    self.conn = cls(self.host, self.port, timeout=self.timeout)
                self.conn.request("POST", self.prefix + path, body=body, headers=headers)
                resp = self.conn.getresponse()
                data = resp.read()  # 必须读净响应体，连接才能复用
                if resp.will_close:
                    self.close()
                return resp.status, data
            except Exception as exc:
                self.close()
                last_exc = exc
        raise last_exc


def maybe_update(resp):
    """自动更新安全阶梯：任一级失败即中止并继续跑旧版。"""
    latest, sha, url = resp.get("latest_version"), resp.get("sha256"), resp.get("update_url")
    if not latest or not sha or not url or latest == __VERSION__:
        return  # 版本号「不同」即更新（允许服务端主动回滚）
    try:  # 冷却：10 分钟内只试一次，防更新风暴
        if os.path.exists(UPDATE_MARKER) and \
                time.time() - os.path.getmtime(UPDATE_MARKER) < UPDATE_COOLDOWN:
            return
        with open(UPDATE_MARKER, "w") as f:  # 下载前先落盘标记
            f.write(str(time.time()))
    except OSError as exc:
        log("更新中止：无法写冷却标记 %s" % exc)
        return
    log("发现新版 %s（当前 %s），开始下载 %s" % (latest, __VERSION__, url))
    try:
        blob = urllib.request.urlopen(url, timeout=30).read()
    except Exception as exc:
        log("更新中止：下载失败 %s" % exc)
        return
    if len(blob) < 4096:
        log("更新中止：文件过小（%d B）" % len(blob))
        return
    if hashlib.sha256(blob).hexdigest() != sha:
        log("更新中止：sha256 不匹配")
        return
    try:
        text = blob.decode("utf-8")
    except UnicodeDecodeError:
        log("更新中止：文件不是 UTF-8")
        return
    m = re.search(r'__VERSION__\s*=\s*"([^"]+)"', text)
    if not m or m.group(1) != latest:
        log("更新中止：下载文件的版本号 %s 与服务端宣称的 %s 不符（下载源缓存陈旧？）"
            % (m.group(1) if m else "?", latest))
        return
    new_path = SELF_PATH + ".new"
    try:
        with open(new_path, "wb") as f:
            f.write(blob)
        py_compile.compile(new_path, cfile=new_path + "c", doraise=True)
        os.remove(new_path + "c")
    except Exception as exc:
        log("更新中止：编译检查失败 %s" % exc)
        try:
            os.remove(new_path)
        except OSError:
            pass
        return
    try:
        shutil.copyfile(SELF_PATH, SELF_PATH + ".prev")  # 手动回滚用
        os.replace(new_path, SELF_PATH)
    except OSError as exc:
        log("更新中止：替换失败 %s" % exc)
        return
    log("已更新到 %s，退出交给 systemd 拉起新版" % latest)
    sys.exit(0)


def main():
    cfg = load_config()
    log("sentinel-agent %s 启动：节点 %s → %s" % (__VERSION__, cfg["name"], cfg["server_url"]))
    link = ServerLink(cfg["server_url"], timeout=10)
    headers = {"Content-Type": "application/json",
               "Authorization": "Bearer " + cfg["token"],
               "User-Agent": "sentinel-agent/" + __VERSION__}
    state = {"prev_stat": None, "prev_net": None, "processes": [],
             "interval": float(cfg.get("interval") or 10)}
    while True:
        report = collect(state, cfg)
        try:
            status, raw = link.post_json("/api/agent/report", report, headers)
        except Exception as exc:
            log("上报失败：%s: %s" % (type(exc).__name__, exc))
            time.sleep(state["interval"])
            continue
        if status == 401:
            log("服务端拒绝口令（401），300 秒后重试")
            time.sleep(300)
            continue
        if status != 200:
            log("上报失败：HTTP %s" % status)
            time.sleep(state["interval"])
            continue
        try:
            resp = json.loads(raw.decode("utf-8", "replace"))
        except ValueError:
            resp = None
        if isinstance(resp, dict):
            try:
                state["interval"] = max(2.0, float(resp.get("interval") or state["interval"]))
            except (TypeError, ValueError):
                pass
            procs = resp.get("processes")
            if isinstance(procs, list):
                state["processes"] = [str(p) for p in procs]
            maybe_update(resp)
        time.sleep(state["interval"])


if __name__ == "__main__":
    main()
