#!/usr/bin/env python3
# PSS YAM agent (Linux) — เก็บ metrics แล้ว POST เข้า /api/metrics/ingest
# ใช้ stdlib ล้วน (ไม่ต้องลง pip) รันเป็น systemd service
import json, os, re, time, socket, subprocess, urllib.request, urllib.error

AGENT_VERSION = "v1.0.0"
CONFIG_PATH = os.environ.get("PSSYAM_CONFIG", "/etc/pss-yam-agent/config.json")

REAL_FS = {"ext2","ext3","ext4","xfs","btrfs","zfs","f2fs","reiserfs","jfs","vfat"}
DEVISH = re.compile(r"(^|[.-])(dev|staging|stg|uat|test|sandbox|demo|beta)([.-]|$)")


def load_config():
    with open(CONFIG_PATH) as f:
        return json.load(f)


def read_cpu():
    with open("/proc/stat") as f:
        parts = f.readline().split()[1:]
    vals = [int(x) for x in parts]
    idle = vals[3] + (vals[4] if len(vals) > 4 else 0)  # idle + iowait
    total = sum(vals)
    return idle, total


def cpu_pct(prev, cur):
    di = cur[0] - prev[0]
    dt = cur[1] - prev[1]
    if dt <= 0:
        return 0.0
    return round((1 - di / dt) * 100, 1)


def mem_mb():
    info = {}
    with open("/proc/meminfo") as f:
        for line in f:
            k, v = line.split(":")
            info[k.strip()] = int(v.strip().split()[0])  # kB
    total = info.get("MemTotal", 0) / 1024
    avail = info.get("MemAvailable", info.get("MemFree", 0)) / 1024
    return round(total), round(total - avail)


def disks():
    out, seen = [], set()
    try:
        with open("/proc/mounts") as f:
            mounts = f.readlines()
    except OSError:
        return out
    for line in mounts:
        p = line.split()
        if len(p) < 3:
            continue
        dev, mount, fstype = p[0], p[1], p[2]
        if fstype not in REAL_FS or mount in seen:
            continue
        seen.add(mount)
        try:
            s = os.statvfs(mount)
            total = s.f_blocks * s.f_frsize
            used = total - s.f_bfree * s.f_frsize
            if total <= 0:
                continue
            out.append({"mount": mount, "total_gb": round(total / 1073741824, 1),
                        "used_gb": round(used / 1073741824, 1)})
        except OSError:
            pass
    return out


def run(cmd, timeout=15):
    return subprocess.run(cmd, capture_output=True, text=True, timeout=timeout).stdout


def services(cfg):
    out = []
    for name in cfg.get("watch_services", []):
        try:
            r = subprocess.run(["systemctl", "is-active", name], capture_output=True, text=True, timeout=8)
            st = r.stdout.strip()
            out.append({"name": name, "status": "running" if st == "active" else st or "unknown"})
        except Exception:
            out.append({"name": name, "status": "unknown"})
    return out


def parse_mem_mb(s):
    m = re.match(r"([0-9.]+)\s*([A-Za-z]+)", s.strip())
    if not m:
        return None
    v = float(m.group(1)); u = m.group(2).lower()
    mult = {"gib": 1024, "mib": 1, "kib": 1/1024, "gb": 1000, "mb": 1, "kb": 1/1000, "b": 1/1048576}
    return round(v * mult.get(u, 1), 1)


def docker():
    out = []
    try:
        txt = run(["docker", "stats", "--no-stream", "--no-trunc",
                   "--format", "{{.Name}}\t{{.CPUPerc}}\t{{.MemUsage}}"], 20)
        for line in txt.splitlines():
            p = line.split("\t")
            if len(p) < 3:
                continue
            cpu = None
            try:
                cpu = round(float(p[1].replace("%", "").strip()), 1)
            except ValueError:
                pass
            ram = parse_mem_mb(p[2].split("/")[0])
            out.append({"name": p[0], "state": "running", "cpu_pct": cpu, "ram_mb": ram})
    except Exception as e:
        log(f"[docker] {e}")
    return out


def postgres(cfg):
    out = []
    for inst in cfg.get("postgres", []):
        try:
            uri = (f"host={inst.get('host','127.0.0.1')} port={inst.get('port',5432)} "
                   f"dbname={inst.get('db','postgres')} user={inst.get('user','')} "
                   f"password={inst.get('password','')} connect_timeout=8")
            conns = run(["psql", uri, "-tAc", "SELECT count(*) FROM pg_stat_activity"], 12).strip()
            out.append({"instance": inst.get("name", "main"), "connections": int(conns) if conns.isdigit() else -1})
        except Exception as e:
            log(f"[pg] {inst.get('name')} {e}")
            out.append({"instance": inst.get("name", "main"), "connections": -1})
    return out


def mysql(cfg):
    out = []
    for inst in cfg.get("mysql", []):
        name = inst.get("name", "main")
        try:
            args = ["mysql", "-N", "-B", "-e", "SHOW STATUS LIKE 'Threads_connected'"]
            if inst.get("host"): args += ["-h", str(inst["host"])]
            if inst.get("port"): args += ["-P", str(inst["port"])]
            if inst.get("user"): args += ["-u", str(inst["user"])]
            env = os.environ.copy()
            if inst.get("password"): env["MYSQL_PWD"] = str(inst["password"])
            r = subprocess.run(args, capture_output=True, text=True, timeout=10, env=env)
            toks = r.stdout.split()
            conns = int(toks[-1]) if toks and toks[-1].isdigit() else -1
            out.append({"instance": name, "connections": conns})
        except Exception as e:
            log(f"[mysql] {name} {e}")
            out.append({"instance": name, "connections": -1})
    return out


def discover_hosts():
    hosts = set()
    # nginx -T dumps full effective config incl. server_name
    try:
        txt = run(["nginx", "-T"], 15)
        for m in re.finditer(r"server_name\s+([^;]+);", txt):
            for h in m.group(1).split():
                hosts.add(h.strip().lower())
    except Exception:
        pass
    # apache vhosts
    for tool in (["apache2ctl", "-S"], ["apachectl", "-S"], ["httpd", "-S"]):
        try:
            txt = run(tool, 15)
            for m in re.finditer(r"(?:namevhost|ServerName)\s+([A-Za-z0-9._-]+)", txt):
                hosts.add(m.group(1).strip().lower())
            break
        except Exception:
            continue
    good = []
    for h in hosts:
        if not h or "." not in h or h in ("_", "localhost"):
            continue
        if h.startswith(("www.", "ipv4.", "webmail.", "*.")):
            continue
        if re.match(r"^\d{1,3}(\.\d{1,3}){3}$", h):
            continue
        if DEVISH.search(h):
            continue
        good.append(h)
    return sorted(set(good))


_net_prev = {"rx": None, "tx": None, "t": None}


def net_mbps():
    try:
        rx = tx = 0
        with open("/proc/net/dev") as f:
            for line in f.readlines()[2:]:
                iface, data = line.split(":", 1)
                if iface.strip() == "lo":
                    continue
                cols = data.split()
                rx += int(cols[0]); tx += int(cols[8])
        now = time.time()
        if _net_prev["t"] is None:
            _net_prev.update(rx=rx, tx=tx, t=now); return (0.0, 0.0)
        dt = now - _net_prev["t"]
        if dt <= 0:
            return (0.0, 0.0)
        r = round(max(0, (rx - _net_prev["rx"]) / 1048576 / dt), 2)
        w = round(max(0, (tx - _net_prev["tx"]) / 1048576 / dt), 2)
        _net_prev.update(rx=rx, tx=tx, t=now)
        return (r, w)
    except Exception:
        return (None, None)


_disk_prev = {"r": None, "w": None, "t": None}
DISK_RE = re.compile(r"^(sd[a-z]+|vd[a-z]+|xvd[a-z]+|nvme\d+n\d+)$")


def disk_io():
    try:
        sr = sw = 0
        with open("/proc/diskstats") as f:
            for line in f:
                p = line.split()
                if len(p) < 10 or not DISK_RE.match(p[2]):
                    continue
                sr += int(p[5]); sw += int(p[9])  # sectors read / written
        now = time.time()
        if _disk_prev["t"] is None:
            _disk_prev.update(r=sr, w=sw, t=now); return (0.0, 0.0)
        dt = now - _disk_prev["t"]
        if dt <= 0:
            return (0.0, 0.0)
        rd = round(max(0, (sr - _disk_prev["r"]) * 512 / 1048576 / dt), 2)
        wr = round(max(0, (sw - _disk_prev["w"]) * 512 / 1048576 / dt), 2)
        _disk_prev.update(r=sr, w=sw, t=now)
        return (rd, wr)
    except Exception:
        return (None, None)


def load_avg():
    try:
        with open("/proc/loadavg") as f:
            return [float(x) for x in f.read().split()[:3]]
    except Exception:
        return None


def top_processes():
    try:
        txt = run(["ps", "-eo", "comm,rss,pcpu", "--no-headers"], 10)
        agg = {}
        for line in txt.splitlines():
            p = line.split(None, 2)
            if len(p) < 3:
                continue
            try:
                rss = int(p[1]); cpu = float(p[2])
            except ValueError:
                continue
            a = agg.setdefault(p[0], {"count": 0, "rss": 0, "cpu": 0.0})
            a["count"] += 1; a["rss"] += rss; a["cpu"] += cpu
        rows = [{"name": k, "count": v["count"], "ram_mb": round(v["rss"] / 1024, 1),
                 "cpu_pct": round(v["cpu"], 1)} for k, v in agg.items()]
        rows.sort(key=lambda x: x["ram_mb"], reverse=True)
        return rows[:8]
    except Exception as e:
        log(f"[proc] {e}")
        return []


def build_payload(cfg, cpu, ram, dsk, pg, my, svc, dk):
    total, used = ram
    host = {"cpu_pct": cpu, "ram_total_mb": total, "ram_used_mb": used, "disks": dsk,
            "top_processes": top_processes()}
    rx, tx = net_mbps()
    if rx is not None:
        host["net_rx_mbps"] = rx; host["net_tx_mbps"] = tx
    dr, dw = disk_io()
    if dr is not None:
        host["disk_read_mbps"] = dr; host["disk_write_mbps"] = dw
    la = load_avg()
    if la:
        host["load_avg"] = la
    return {
        "machine_id": cfg["machine_id"],
        "agent_version": AGENT_VERSION,
        "name": cfg.get("host_name", cfg["machine_id"]),
        "ts": time.strftime("%Y-%m-%dT%H:%M:%S", time.gmtime()) + "Z",
        "host": host,
        "postgres": pg, "mysql": my, "services": svc, "docker": dk,
    }


def post(url, token, obj, timeout=20):
    data = json.dumps(obj).encode()
    req = urllib.request.Request(url, data=data, method="POST",
                                 headers={"Content-Type": "application/json",
                                          "Authorization": f"Bearer {token}"})
    with urllib.request.urlopen(req, timeout=timeout) as r:
        return r.status, r.read().decode()


def log(msg):
    print(f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] {msg}", flush=True)


def main():
    log(f"PSS YAM Linux agent {AGENT_VERSION} starting")
    prev = read_cpu(); time.sleep(1)
    last_discover = 0.0
    while True:
        try:
            cfg = load_config()
            base = cfg["backend_url"].rstrip("/")
            token = cfg["agent_token"]
            cur = read_cpu()
            dk = docker() if cfg.get("docker_enabled", True) else []
            payload = build_payload(cfg, cpu_pct(prev, cur), mem_mb(), disks(),
                                    postgres(cfg), mysql(cfg), services(cfg), dk)
            prev = cur
            try:
                st, _ = post(base + "/api/metrics/ingest", token, payload)
                log(f"sent · cpu {payload['host']['cpu_pct']}% ram {payload['host']['ram_used_mb']}/{payload['host']['ram_total_mb']}MB svc {len(payload['services'])} docker {len(payload['docker'])} (HTTP {st})")
            except urllib.error.HTTPError as e:
                log(f"[ingest] HTTP {e.code} {e.read().decode()[:200]}")
            except Exception as e:
                log(f"[ingest] {e}")

            if cfg.get("discover_ssl") and (time.time() - last_discover) >= cfg.get("discover_interval_minutes", 360) * 60:
                last_discover = time.time()
                hosts = discover_hosts()
                if hosts:
                    try:
                        st, body = post(base + "/api/monitors/discover", token,
                                        {"machine_id": cfg["machine_id"], "hosts": hosts})
                        added = json.loads(body).get("added", "?")
                        log(f"[ssl-discover] เจอ {len(hosts)} โฮสต์ · เพิ่มใหม่ {added}")
                    except Exception as e:
                        log(f"[ssl-discover] {e}")
        except Exception as e:
            log(f"[loop] {e}")
        time.sleep(max(5, load_safe_interval()))


def load_safe_interval():
    try:
        with open(CONFIG_PATH) as f:
            return int(json.load(f).get("interval_seconds", 30))
    except Exception:
        return 30


if __name__ == "__main__":
    main()
