#!/usr/bin/env python3
"""Stream PC performance stats to the handheld over USB serial + web dashboard.

Protocol (115200 baud, one line per update):
  P c=45 r=72 g=38 v=55 d=12 ct=52 gt=61 u=1200 dn=5600

Install:
  pip install pyserial psutil

Usage:
  python pc_monitor.py
  python pc_monitor.py --port /dev/ttyACM0
  python pc_monitor.py --web-port 8765
  python pc_monitor.py --push https://dharkangel.com/pc-monitor/api/push.php
"""

from __future__ import annotations

import argparse
import ftplib
import io
import json
import os
import queue
import re
import shutil
import socket
import ssl
import subprocess
import sys
import threading
import time
import urllib.error
import urllib.request
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path

try:
    import psutil
except ImportError:
    print("Missing psutil — run: pip install psutil", file=sys.stderr)
    raise SystemExit(1)

try:
    import serial
    import serial.tools.list_ports
except ImportError:
    print("Missing pyserial — run: pip install pyserial", file=sys.stderr)
    raise SystemExit(1)

class StatsState:
    def __init__(self) -> None:
        self._lock = threading.Lock()
        self._data: dict = {
            "host": socket.gethostname(),
            "serial_port": "",
            "serial_ok": False,
            "updated": 0.0,
            "cpu": -1,
            "ram": -1,
            "gpu": -1,
            "vram": -1,
            "disk": -1,
            "cpu_temp": -1,
            "gpu_temp": -1,
            "net_up": 0,
            "net_dn": 0,
        }

    def update(self, stats: dict[str, int], *, port: str, serial_ok: bool) -> None:
        with self._lock:
            self._data.update(
                {
                    "host": socket.gethostname(),
                    "serial_port": port,
                    "serial_ok": serial_ok,
                    "updated": time.time(),
                    "cpu": stats["c"],
                    "ram": stats["r"],
                    "gpu": stats["g"],
                    "vram": stats["v"],
                    "disk": stats["d"],
                    "cpu_temp": stats["ct"],
                    "gpu_temp": stats["gt"],
                    "net_up": stats["u"],
                    "net_dn": stats["dn"],
                }
            )

    def snapshot(self) -> dict:
        with self._lock:
            return dict(self._data)


def find_esp_port() -> str | None:
    for port in serial.tools.list_ports.comports():
        desc = (port.description or "").lower()
        manu = (port.manufacturer or "").lower()
        if any(k in desc for k in ("usb", "serial", "jtag", "esp", "cdc")):
            return port.device
        if "espressif" in manu or "silicon labs" in manu:
            return port.device
    return None


HERE = Path(__file__).resolve().parent
GPU_REFRESH_S = 3.0
CLOUD_PUSH_S = 5.0
SERIAL_REOPEN_S = 2.0
LOG_MAX_LINES = 20000
LOG_HEADER = (
    "iso_time,unix_time,host,cpu,ram,gpu,vram,disk,"
    "cpu_temp,gpu_temp,net_up,net_dn,serial_ok\n"
)


class CsvLogger:
    def __init__(self, host_id: str) -> None:
        self.host_id = host_id
        self._lock = threading.Lock()
        self._path = (
            Path.home() / ".local" / "share" / "pc-monitor" / "logs" / f"{host_id}.csv"
        )
        self._path.parent.mkdir(parents=True, exist_ok=True)
        if not self._path.exists() or self._path.stat().st_size == 0:
            self._path.write_text(LOG_HEADER, encoding="utf-8")

    @property
    def path(self) -> Path:
        return self._path

    def append(self, snap: dict) -> None:
        updated = float(snap.get("updated") or time.time())
        iso = time.strftime("%Y-%m-%dT%H:%M:%S", time.localtime(updated))
        serial_ok = 1 if snap.get("serial_ok") else 0
        line = (
            f"{iso},{updated:.3f},{snap.get('host', self.host_id)},"
            f"{snap.get('cpu', -1)},{snap.get('ram', -1)},"
            f"{snap.get('gpu', -1)},{snap.get('vram', -1)},"
            f"{snap.get('disk', -1)},{snap.get('cpu_temp', -1)},"
            f"{snap.get('gpu_temp', -1)},"
            f"{snap.get('net_up', 0)},{snap.get('net_dn', 0)},"
            f"{serial_ok}\n"
        )
        with self._lock:
            with self._path.open("a", encoding="utf-8") as handle:
                handle.write(line)
            self._trim_locked()

    def _trim_locked(self) -> None:
        lines = self._path.read_text(encoding="utf-8").splitlines()
        if len(lines) <= LOG_MAX_LINES:
            return
        kept = [lines[0]] + lines[-(LOG_MAX_LINES - 1) :]
        self._path.write_text("\n".join(kept) + "\n", encoding="utf-8")

    def read_bytes(self) -> bytes:
        with self._lock:
            return self._path.read_bytes()


class GpuCache:
    def __init__(self) -> None:
        self._lock = threading.Lock()
        self._stats = (-1, -1, -1)
        self._thread = threading.Thread(target=self._run, daemon=True)
        self._thread.start()

    def _run(self) -> None:
        while True:
            stats = _gpu_stats_once()
            with self._lock:
                self._stats = stats
            time.sleep(GPU_REFRESH_S)

    def get(self) -> tuple[int, int, int]:
        with self._lock:
            return self._stats


def _gpu_stats_once() -> tuple[int, int, int]:
    if shutil.which("nvidia-smi"):
        try:
            out = subprocess.check_output(
                [
                    "nvidia-smi",
                    "--query-gpu=utilization.gpu,utilization.memory,temperature.gpu",
                    "--format=csv,noheader,nounits",
                ],
                text=True,
                timeout=2,
            ).strip()
            parts = [p.strip() for p in out.split(",")]
            if len(parts) >= 3:
                return int(float(parts[0])), int(float(parts[1])), int(float(parts[2]))
        except (subprocess.SubprocessError, ValueError):
            pass
    return -1, -1, -1


def cpu_temp_c() -> int:
    temps = getattr(psutil, "sensors_temperatures", lambda: {})()
    for key in ("coretemp", "k10temp", "cpu_thermal", "acpitz"):
        entries = temps.get(key)
        if entries:
            return int(max(e.current for e in entries if e.current is not None))
    return -1


def sample(
    prev_net: tuple[int, int],
    interval: float,
    gpu_cache: GpuCache,
) -> tuple[dict[str, int], tuple[int, int]]:
    gpu, vram, gt = gpu_cache.get()
    net = psutil.net_io_counters()
    now = (net.bytes_sent, net.bytes_recv)
    up_kb = max(0, int((now[0] - prev_net[0]) / interval / 1024))
    dn_kb = max(0, int((now[1] - prev_net[1]) / interval / 1024))
    disk = psutil.disk_usage("/")
    return (
        {
            "c": int(psutil.cpu_percent(interval=None)),
            "r": int(psutil.virtual_memory().percent),
            "g": gpu,
            "v": vram,
            "d": int(disk.percent),
            "ct": cpu_temp_c(),
            "gt": gt,
            "u": up_kb,
            "dn": dn_kb,
        },
        now,
    )


def format_line(stats: dict[str, int]) -> str:
    return (
        f"P c={stats['c']} r={stats['r']} g={stats['g']} v={stats['v']} "
        f"d={stats['d']} ct={stats['ct']} gt={stats['gt']} "
        f"u={stats['u']} dn={stats['dn']}\n"
    )


def push_stats_http(url: str, payload: dict) -> None:
    body = json.dumps(payload).encode("utf-8")
    req = urllib.request.Request(
        url,
        data=body,
        headers={
            "Content-Type": "application/json",
            "User-Agent": "PCMonitor/1.0",
        },
        method="POST",
    )
    with urllib.request.urlopen(req, timeout=4) as resp:
        resp.read()


def load_cloud_config() -> dict | None:
    cfg_path = Path.home() / ".config" / "pc-monitor" / "cloud.json"
    if cfg_path.exists():
        return json.loads(cfg_path.read_text(encoding="utf-8"))

    user = os.environ.get("PC_MONITOR_FTP_USER", "")
    password = os.environ.get("PC_MONITOR_FTP_PASS", "")
    if user and password:
        return {
            "ftp_host": os.environ.get("PC_MONITOR_FTP_HOST", "ftp.jez.thu.mybluehost.me"),
            "ftp_user": user,
            "ftp_pass": password,
            "ftp_path": os.environ.get(
                "PC_MONITOR_FTP_PATH", "/public_html/pc-monitor/data"
            ),
        }
    return None


def cloud_host_id() -> str:
    host = re.sub(r"[^a-zA-Z0-9_-]", "", socket.gethostname())
    return host or "unknown"


def push_stats_ftp(payload: dict, cfg: dict) -> None:
    host = cloud_host_id()
    payload = dict(payload)
    payload["host"] = host
    body = json.dumps(payload).encode("utf-8")

    ctx = ssl.create_default_context()
    ctx.check_hostname = False
    ctx.verify_mode = ssl.CERT_NONE

    data_path = cfg.get("ftp_path", "/public_html/pc-monitor/data")
    ftp = ftplib.FTP_TLS(context=ctx)
    ftp.connect(cfg.get("ftp_host", "ftp.jez.thu.mybluehost.me"), 21, timeout=20)
    ftp.login(cfg["ftp_user"], cfg["ftp_pass"])
    ftp.prot_p()
    ftp.cwd(data_path)
    ftp.storbinary(f"STOR {host}.json", io.BytesIO(body))
    ftp.storbinary("STOR latest.json", io.BytesIO(body))

    skip = {"hosts", "latest", "config-test", "testhost"}
    names = {host}
    for name in ftp.nlst():
        if name.endswith(".json") and name[:-5] not in skip:
            names.add(name[:-5])
    ftp.storbinary("STOR hosts.json", io.BytesIO(json.dumps(sorted(names)).encode()))
    log_path = Path.home() / ".local" / "share" / "pc-monitor" / "logs" / f"{host}.csv"
    if log_path.exists():
        ftp.storbinary(f"STOR {host}.log.csv", io.BytesIO(log_path.read_bytes()))
    ftp.quit()


def make_handler(state: StatsState, live_html: bytes, csv_logger: CsvLogger | None):
    class Handler(BaseHTTPRequestHandler):
        def log_message(self, fmt: str, *args) -> None:
            return

        def _cors(self) -> None:
            self.send_header("Access-Control-Allow-Origin", "*")
            self.send_header("Access-Control-Allow-Methods", "GET, OPTIONS")
            self.send_header("Cache-Control", "no-store")

        def do_OPTIONS(self) -> None:
            self.send_response(204)
            self._cors()
            self.end_headers()

        def do_GET(self) -> None:
            if self.path in ("/", "/live", "/live.html"):
                self.send_response(200)
                self._cors()
                self.send_header("Content-Type", "text/html; charset=utf-8")
                self.end_headers()
                self.wfile.write(live_html)
                return

            if self.path == "/api/stats":
                data = json.dumps(state.snapshot()).encode("utf-8")
                self.send_response(200)
                self._cors()
                self.send_header("Content-Type", "application/json")
                self.end_headers()
                self.wfile.write(data)
                return

            if self.path == "/api/log.csv" and csv_logger is not None:
                data = csv_logger.read_bytes()
                self.send_response(200)
                self._cors()
                self.send_header("Content-Type", "text/csv; charset=utf-8")
                self.send_header(
                    "Content-Disposition",
                    f'attachment; filename="{csv_logger.host_id}.log.csv"',
                )
                self.end_headers()
                self.wfile.write(data)
                return

            self.send_response(404)
            self.end_headers()

    return Handler


def start_web_server(
    state: StatsState,
    port: int,
    csv_logger: CsvLogger | None,
) -> tuple[ThreadingHTTPServer, int]:
    live_path = HERE / "live.html"
    live_html = live_path.read_bytes() if live_path.exists() else b"<h1>live.html missing</h1>"
    last_err: OSError | None = None
    for attempt in range(10):
        try_port = port + attempt
        try:
            server = ThreadingHTTPServer(
                ("127.0.0.1", try_port),
                make_handler(state, live_html, csv_logger),
            )
            server.daemon_threads = True
            thread = threading.Thread(target=server.serve_forever, daemon=True)
            thread.start()
            if attempt:
                print(f"Port {port} busy — dashboard on {try_port}")
            return server, try_port
        except OSError as exc:
            if exc.errno not in (98, 48):  # Linux / macOS address in use
                raise
            last_err = exc
    raise SystemExit(f"No free port near {port}: {last_err}")


def open_serial(port: str, baud: int) -> serial.Serial:
    ser = serial.Serial(port, baud, timeout=0.1, write_timeout=0.5)
    # Avoid toggling DTR/RTS — some ESP32 boards reset when the port opens.
    ser.dtr = False
    ser.rts = False
    return ser


def cloud_push_loop(cfg: dict, q: queue.Queue[dict]) -> None:
    while True:
        payload = q.get()
        while True:
            try:
                latest = q.get_nowait()
            except queue.Empty:
                break
            payload = latest
        try:
            push_stats_ftp(payload, cfg)
        except (OSError, ftplib.Error) as exc:
            print(f"Cloud push failed: {exc}", file=sys.stderr)


def main() -> None:
    parser = argparse.ArgumentParser(description="USB PC stats for handheld monitor")
    parser.add_argument("--port", help="Serial port (auto-detect if omitted)")
    parser.add_argument("--baud", type=int, default=115200)
    parser.add_argument("--hz", type=float, default=2.0, help="Update rate")
    parser.add_argument("--web-port", type=int, default=8765, help="Local dashboard port (0=off)")
    parser.add_argument(
        "--cloud",
        action="store_true",
        help="Upload stats via FTP (see ~/.config/pc-monitor/cloud.json)",
    )
    parser.add_argument(
        "--push",
        help="Optional HTTP POST URL (often blocked by Cloudflare — prefer --cloud)",
    )
    parser.add_argument("--no-serial", action="store_true", help="Web/push only, skip handheld")
    args = parser.parse_args()

    serial_port = args.port or find_esp_port()
    if not args.no_serial and not serial_port:
        print("No serial port found — pass --port or use --no-serial", file=sys.stderr)
        raise SystemExit(1)

    interval = 1.0 / max(args.hz, 0.2)
    state = StatsState()
    gpu_cache = GpuCache()
    csv_logger = CsvLogger(cloud_host_id())
    cloud_queue: queue.Queue[dict] | None = None

    if args.web_port:
        _server, web_port = start_web_server(state, args.web_port, csv_logger)
        print(f"Live dashboard: http://127.0.0.1:{web_port}/")
        print(f"Session log: {csv_logger.path}")

    cloud_cfg = load_cloud_config() if args.cloud else None
    if args.cloud and not cloud_cfg:
        print(
            "Cloud push needs ~/.config/pc-monitor/cloud.json or PC_MONITOR_FTP_* env vars",
            file=sys.stderr,
        )
        raise SystemExit(1)
    if cloud_cfg:
        cid = cloud_host_id()
        print(f"Cloud push: FTP → dharkangel.com (every {CLOUD_PUSH_S:.0f}s, background)")
        print(f"Cloud host ID: {cid}")
        print(f"Live URL: https://dharkangel.com/pc-monitor/live.html?host={cid}")
        cloud_queue = queue.Queue(maxsize=1)
        threading.Thread(
            target=cloud_push_loop,
            args=(cloud_cfg, cloud_queue),
            daemon=True,
        ).start()
    elif args.push:
        print(f"Cloud push: HTTP {args.push}")

    if serial_port and not args.no_serial:
        print(f"Serial: {serial_port} @ {args.baud} ({args.hz:.1f} Hz)")
    print("Open PC Monitor on the handheld, then leave this running.")
    print("Ctrl+C to stop.\n")

    ser: serial.Serial | None = None
    if serial_port and not args.no_serial:
        ser = open_serial(serial_port, args.baud)
        time.sleep(0.3)

    prev_net = (0, 0)
    psutil.cpu_percent(interval=None)
    tick = 0
    push_errors = 0
    last_cloud_push = 0.0
    last_serial_open = time.monotonic()

    try:
        while True:
            t0 = time.monotonic()
            stats, prev_net = sample(prev_net, interval, gpu_cache)
            serial_ok = False

            if ser is not None:
                try:
                    ser.write(format_line(stats).encode("ascii"))
                    ser.flush()
                    serial_ok = True
                except serial.SerialException as exc:
                    serial_ok = False
                    if time.monotonic() - last_serial_open >= SERIAL_REOPEN_S:
                        last_serial_open = time.monotonic()
                        try:
                            ser.close()
                        except serial.SerialException:
                            pass
                        try:
                            ser = open_serial(serial_port, args.baud)
                            print(f"Serial reopened: {serial_port}", file=sys.stderr)
                        except serial.SerialException as reopen_exc:
                            ser = None
                            print(f"Serial lost ({exc}); reopen failed: {reopen_exc}", file=sys.stderr)

            state.update(stats, port=serial_port or "", serial_ok=serial_ok)
            snap = state.snapshot()
            csv_logger.append(snap)

            now = time.time()
            if cloud_queue is not None and (now - last_cloud_push) >= CLOUD_PUSH_S:
                last_cloud_push = now
                try:
                    cloud_queue.put_nowait(snap)
                except queue.Full:
                    try:
                        cloud_queue.get_nowait()
                    except queue.Empty:
                        pass
                    cloud_queue.put_nowait(snap)
            elif args.push and (now - last_cloud_push) >= CLOUD_PUSH_S:
                last_cloud_push = now
                try:
                    push_stats_http(args.push, snap)
                    push_errors = 0
                except (urllib.error.URLError, TimeoutError, OSError) as exc:
                    push_errors += 1
                    if push_errors == 1 or push_errors % 20 == 0:
                        print(f"Cloud push failed: {exc}", file=sys.stderr)

            tick += 1
            if tick <= 3 or tick % 10 == 0:
                print(format_line(stats).strip())

            elapsed = time.monotonic() - t0
            time.sleep(max(0.0, interval - elapsed))
    except KeyboardInterrupt:
        print("\nStopped.")
    finally:
        if ser is not None:
            ser.close()


if __name__ == "__main__":
    main()
