#!/usr/bin/env python3
"""UI-only Cyber Protect guardian — grounded answers from live box state.

Never blocks devices, never mentions Factory/prices, never changes network policy.
Factual network answers always come from local state. Optional Atlas LLM only
rewrites open-ended chat using the same facts (never invents devices/threats).
"""
from __future__ import annotations

import json
import os
import re
import ssl
import time
import urllib.request
from pathlib import Path
from typing import Any

DATA = Path("/var/lib/atlas-cyber-protect")
STATE = DATA / "state.json"
ATLAS_URL = os.environ.get("ATLAS_URL", "https://atlas-server.taile9cc75.ts.net").rstrip("/")


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

        return _reach()
    except Exception:
        return ATLAS_URL


def load_state() -> dict[str, Any]:
    if STATE.is_file():
        try:
            return json.loads(STATE.read_text(encoding="utf-8"))
        except Exception:
            pass
    return {}


def save_state(st: dict[str, Any]) -> None:
    """Merge-save so portal/guardian never wipe agent keys like seen_macs."""
    DATA.mkdir(parents=True, exist_ok=True)
    cur: dict[str, Any] = {}
    if STATE.is_file():
        try:
            cur = json.loads(STATE.read_text(encoding="utf-8"))
            if not isinstance(cur, dict):
                cur = {}
        except Exception:
            cur = {}
    cur.update(st)
    # Preserve critical agent maps if caller omitted them
    for k in ("seen_macs", "traffic_totals", "timeout_counts", "decisions", "last_devices"):
        if k not in st and k in cur:
            pass  # already in cur
        st.setdefault(k, cur.get(k) if k in cur else ({} if k != "last_devices" else []))
        cur[k] = st.get(k, cur.get(k))
    STATE.write_text(json.dumps(cur), encoding="utf-8")


def _devices(st: dict) -> list[dict]:
    d = st.get("last_devices") or []
    return d if isinstance(d, list) else []


def _open_findings(st: dict) -> list[dict]:
    decided = st.get("decisions") or {}
    return [f for f in (st.get("open_findings") or []) if (f.get("key") or "") not in decided]


def _threats(st: dict) -> list[dict]:
    return [
        f
        for f in _open_findings(st)
        if f.get("category") == "threat" or f.get("severity") in ("critical", "high", "medium")
    ]


def security_score(st: dict) -> dict[str, Any]:
    devices = _devices(st)
    n = len(devices)
    up = sum(1 for d in devices if d.get("reachable"))
    inet = bool(st.get("internet_ok", True))
    learn = float(st.get("learning_pct") or 0)
    findings = _open_findings(st)
    open_n = len(findings)
    threats = _threats(st)
    crit = sum(1 for f in threats if f.get("severity") == "critical")

    score = 40
    if inet:
        score += 20
    score += min(20, int(learn / 5))
    if n >= 3:
        score += 10
    if up and n:
        score += min(10, int(10 * up / max(1, n)))
    score -= min(35, open_n * 4 + crit * 8)
    score = max(5, min(99, score))

    if not inet:
        label, tone = "Internet trouble", "bad"
    elif crit:
        label, tone = "Critical threats open", "bad"
    elif threats:
        label, tone = "Threats need decisions", "warn"
    elif open_n:
        label, tone = "Watching closely", "warn"
    elif learn < 100:
        label, tone = "Learning your network", "sky"
    else:
        label, tone = "Shields strong", "ok"

    return {
        "score": score,
        "label": label,
        "tone": tone,
        "devices": n,
        "reachable": up,
        "open_issues": open_n,
        "threats": len(threats),
        "critical": crit,
        "learning_pct": learn,
        "internet_ok": inet,
    }


def shields(st: dict) -> list[dict[str, Any]]:
    devices = _devices(st)
    sm = st.get("scan_meta") or {}
    up = sum(1 for d in devices if d.get("reachable"))
    slow = sum(
        1
        for d in devices
        if d.get("reachable") and isinstance(d.get("ping_ms"), (int, float)) and float(d["ping_ms"]) >= 80
    )
    timed = sum(1 for d in devices if d.get("ip") and not d.get("reachable"))
    talkers = sorted(devices, key=lambda x: int(x.get("bytes_obs") or 0), reverse=True)
    top = talkers[0] if talkers and int(talkers[0].get("bytes_obs") or 0) > 0 else None
    open_threats = _threats(st)
    ts = st.get("threat_summary") or {}
    threat_status = "ok"
    if any(f.get("severity") == "critical" for f in open_threats):
        threat_status = "bad"
    elif open_threats:
        threat_status = "warn"
    threat_detail = (
        f"{ts.get('findings', 0)} open · {ts.get('critical', 0)} critical · {ts.get('high', 0)} high — open Threats tab"
        if ts.get("checked_at")
        else "Scanning risky ports, new devices, gateway integrity…"
    )
    return [
        {"id": "threats", "title": "Threat hunt", "status": threat_status, "detail": threat_detail},
        {
            "id": "internet",
            "title": "Internet path",
            "status": "ok" if st.get("internet_ok", True) else "bad",
            "detail": "Gateway path answering" if st.get("internet_ok", True) else "No route to the internet from this box",
        },
        {
            "id": "lan",
            "title": "LAN coverage",
            "status": "ok" if devices else "warn",
            "detail": f"{len(devices)} devices on {sm.get('cidr') or 'your LAN'} · {up} answering",
        },
        {
            "id": "latency",
            "title": "Latency watch",
            "status": "warn" if slow else "ok",
            "detail": f"{slow} slow · {timed} timeout" if (slow or timed) else "Response times look healthy",
        },
        {
            "id": "bandwidth",
            "title": "Bandwidth pulse",
            "status": "warn" if top and float(top.get("traffic_share_pct") or 0) >= 40 else "ok",
            "detail": (
                f"Top talker: {top.get('hostname') or top.get('display_name') or top.get('ip')} (~{top.get('traffic_share_pct')}%)"
                if top
                else "Sampling live traffic visible to Protect"
            ),
        },
        {
            "id": "learning",
            "title": "Behaviour model",
            "status": "ok" if float(st.get("learning_pct") or 0) >= 100 else "sky",
            "detail": f"Learning {float(st.get('learning_pct') or 0):.0f}% — normal patterns for this site",
        },
    ]


def live_feed(st: dict, *, limit: int = 14) -> list[dict[str, Any]]:
    feed = list(st.get("guardian_feed") or [])
    devices = _devices(st)
    now = time.time()
    sm = st.get("scan_meta") or {}
    pulse = security_score(st)
    ts = st.get("threat_summary") or {}

    lines = [
        f"Guardian pulse {pulse['score']}/100 — {pulse['label']}.",
        f"Full LAN sweep on {sm.get('cidr') or 'your network'}: {len(devices)} devices, {pulse['reachable']} answering.",
    ]
    if ts.get("checked_at"):
        lines.append(
            f"Threat hunt: {ts.get('findings', 0)} open ({ts.get('critical', 0)} critical, {ts.get('high', 0)} high)."
        )
    if st.get("internet_ok", True):
        lines.append("Internet path from Protect is healthy.")
    else:
        lines.append("Internet looks down from this box — check the router or backup link.")

    talkers = sorted(devices, key=lambda x: int(x.get("bytes_obs") or 0), reverse=True)[:3]
    for t in talkers:
        if int(t.get("bytes_obs") or 0) <= 0:
            continue
        lines.append(
            f"Bandwidth: {t.get('display_name') or t.get('hostname') or t.get('ip') or t.get('mac')} "
            f"~{t.get('traffic_share_pct')}% observed · {t.get('bytes_rate') or 0} B/s."
        )

    for f in _threats(st)[:5]:
        lines.append(f"Threat: [{f.get('severity')}] {f.get('title') or f.get('plain_tip')}")

    lines.append("I never silently block — you choose Fine or Neutralize on the Threats tab.")

    existing_text = {str(x.get("text") or "") for x in feed[:40]}
    for text in reversed(lines):
        if text in existing_text:
            continue
        feed.insert(0, {"ts": now, "text": text, "kind": "pulse"})
        existing_text.add(text)
    feed = feed[:60]
    st["guardian_feed"] = feed
    st["guardian_pulse"] = pulse
    st["guardian_shields"] = shields(st)
    st["guardian_updated_at"] = now
    return feed[:limit]


def _facts_pack(st: dict) -> dict[str, Any]:
    """Compact truth blob every answer must respect."""
    devices = _devices(st)
    threats = _threats(st)
    nh = st.get("network_health") or {}
    ts = st.get("threat_summary") or {}
    pulse = security_score(st)
    talkers = sorted(devices, key=lambda x: int(x.get("bytes_obs") or 0), reverse=True)[:5]
    return {
        "pulse": pulse,
        "scan_meta": st.get("scan_meta") or {},
        "threat_summary": ts,
        "threats": [
            {
                "severity": f.get("severity"),
                "title": f.get("title"),
                "plain": f.get("plain_tip"),
                "key": f.get("key"),
            }
            for f in threats[:12]
        ],
        "devices": [
            {
                "ip": d.get("ip"),
                "mac": d.get("mac"),
                "name": d.get("display_name") or d.get("hostname") or d.get("vendor") or d.get("ip"),
                "vendor": d.get("vendor"),
                "type": d.get("device_type"),
                "icon": d.get("icon_key"),
                "ping_ms": d.get("ping_ms"),
                "reachable": d.get("reachable"),
                "ports": d.get("open_ports"),
                "traffic_pct": d.get("traffic_share_pct"),
            }
            for d in devices[:30]
        ],
        "talkers": [
            {
                "name": t.get("display_name") or t.get("hostname") or t.get("ip"),
                "pct": t.get("traffic_share_pct"),
                "rate": t.get("bytes_rate"),
            }
            for t in talkers
            if int(t.get("bytes_obs") or 0) > 0
        ],
        "services": (st.get("services_summary") or [])[:12],
        "network_health": {
            "internet_ok": nh.get("internet_ok"),
            "dns_ok": (nh.get("dns") or {}).get("ok"),
            "gateway": (nh.get("gateway") or {}).get("gateway"),
            "targets": nh.get("targets") or [],
        },
        "capabilities_active": [c.get("title") for c in (st.get("capabilities") or []) if c.get("active")],
        "ui_tabs": ["Overview", "Devices", "Traffic", "Insights", "Threats", "Ask guardian", "How it works"],
    }


def _find_device(q: str, devices: list[dict]) -> dict | None:
    ql = q.lower()
    best = None
    for d in devices:
        needles = [
            str(d.get("ip") or ""),
            str(d.get("hostname") or ""),
            str(d.get("display_name") or ""),
            str(d.get("vendor") or ""),
            str(d.get("mac") or ""),
        ]
        for n in needles:
            n = n.lower().strip()
            if len(n) >= 3 and n in ql:
                return d
        # partial hostname token
        host = str(d.get("hostname") or "").lower()
        if host and any(tok and tok in host for tok in re.findall(r"[a-z0-9-]{4,}", ql)):
            best = best or d
    return best


def _local_answer(question: str, st: dict) -> tuple[str, str]:
    """Return (reply, intent). Intent 'open' means Atlas may polish."""
    q = (question or "").strip()
    ql = q.lower()
    devices = _devices(st)
    pulse = security_score(st)
    sm = st.get("scan_meta") or {}
    facts = _facts_pack(st)
    threats = _threats(st)
    ts = st.get("threat_summary") or {}

    if not ql:
        return (
            f"I'm your Cyber Protect guardian. Pulse {pulse['score']}/100 — {pulse['label']}. "
            f"{len(devices)} devices on {sm.get('cidr') or 'your LAN'}, {len(threats)} open threat(s). "
            "Ask about threats, a device, bandwidth, or any UI tab.",
            "status",
        )

    # UI tab help
    if any(w in ql for w in ("threats tab", "decisions tab", "explain the threat")):
        return (
            "Threats tab lists LAN risks Protect found (risky ports like Telnet/SMB/RDP, new devices, "
            "gateway MAC changes, inbound probes). Each item has Fine (trust for now) or Neutralize "
            f"(contain that MAC on Protect’s path). Right now: {ts.get('findings', len(threats))} open, "
            f"{ts.get('critical', 0)} critical.",
            "ui",
        )
    if "how it works" in ql or "capabilities" in ql:
        active = facts.get("capabilities_active") or []
        return (
            "How it works explains Protect as a silent LAN side-box and shows live capabilities with proof. "
            + (f"Active now: {'; '.join(active[:8])}." if active else "Capabilities warm up after the first scan."),
            "ui",
        )
    if any(w in ql for w in ("devices tab", "traffic tab", "insights tab", "overview tab")):
        return (
            "Overview = pulse + shields + live processes. Devices = inventory with icons, ping, ports. "
            "Traffic = bandwidth talkers + services. Insights = internet/DNS + device timeline + PC agents. "
            "Threats = security findings. Ask guardian = me. How it works = proof of what this box does.",
            "ui",
        )

    # PC agents / antivirus / programs
    raw_agents = st.get("companion_agents") or {}
    if isinstance(raw_agents, dict):
        agents = list(raw_agents.values())
    elif isinstance(raw_agents, list):
        agents = raw_agents
    else:
        agents = []
    if any(w in ql for w in ("antivirus", "anti virus", "defender", "pc agent", "windows agent", "virus scan", "av on")):
        if not agents:
            return (
                "No PC agents paired yet. Install Atlas Cyber Protect Agent on each Windows PC. "
                "The box watches the LAN; the app uses Windows Defender or another installed antivirus to monitor/scan. "
                "Atlas never deletes files — the AV engine does its own job.",
                "endpoint",
            )
        bits = []
        for a in agents[:8]:
            inv = a.get("inventory") if isinstance(a.get("inventory"), dict) else {}
            av = inv.get("antivirus") if isinstance(inv.get("antivirus"), dict) else {}
            primary = av.get("primary") if isinstance(av.get("primary"), dict) else {}
            bits.append(
                f"{a.get('name') or a.get('agent_id')}: "
                f"{primary.get('name') or inv.get('primary_av') or 'AV unknown'} "
                f"({primary.get('engine') or inv.get('av_engine') or '?'})"
            )
        return (
            "PC antivirus (from paired agents): " + "; ".join(bits) + ". "
            "Queue a scan from Insights → PC Agents, or use AV Scan on the PC app. "
            "Atlas learns these reports so I know what’s on each computer.",
            "endpoint",
        )
    if any(w in ql for w in ("program", "installed app", "software on", "what apps", "applications on")):
        if not agents:
            return (
                "No PC software inventory yet — pair a Windows Protect Agent first. "
                "It reports installed programs to this box on each heartbeat.",
                "endpoint",
            )
        lines = []
        for a in agents[:5]:
            inv = a.get("inventory") if isinstance(a.get("inventory"), dict) else {}
            software = inv.get("software") if isinstance(inv.get("software"), dict) else {}
            progs = software.get("programs") if isinstance(software.get("programs"), list) else []
            names = [str(p.get("name") or "")[:40] for p in progs[:6] if isinstance(p, dict) and p.get("name")]
            lines.append(
                f"{a.get('name') or 'PC'}: {int(software.get('count') or len(progs) or 0)} programs"
                + (f" (e.g. {', '.join(names)})" if names else "")
            )
        return ("Installed programs from PC agents: " + " | ".join(lines), "endpoint")

    if any(w in ql for w in ("email link", "phishing", "dangerous link", "hacker link", "link warning", "clipboard link")):
        hits = list(st.get("link_guard_hits") or [])
        if not hits:
            try:
                import link_guard as lg

                hits = list(lg.load_hits(12))
            except Exception:
                hits = []
        if not hits:
            return (
                "No email/link warnings logged yet. On each PC the Protect Agent watches clipboard links "
                "and can scan Outlook/email text. It warns if a link matches phishing/malware host data "
                "or looks like a classic email trick — it does not block your internet.",
                "endpoint",
            )
        bits = []
        for h in hits[:8]:
            bits.append(
                f"{int(h.get('danger') or 0)} danger / {int(h.get('warn') or 0)} warn"
                + (f" ({', '.join((h.get('hosts') or [])[:3])})" if h.get("hosts") else "")
            )
        return (
            "Recent email/link warnings from PC agents: " + "; ".join(bits) + ". "
            "Tell staff not to open those links. Protect learns each hit for this site.",
            "endpoint",
        )

    # Threats
    if any(w in ql for w in ("threat", "virus", "malware", "attack", "risk", "ransomware", "telnet", "rdp", "smb")):
        if "telnet" in ql or "rdp" in ql or "smb" in ql or "ftp" in ql:
            port_names = {"telnet": "23", "rdp": "3389", "smb": "445", "ftp": "21"}
            hit = [name for name in port_names if name in ql]
            matched = []
            for f in threats:
                title = (f.get("title") or "").lower()
                if any(h in title for h in hit) or any(h in (f.get("plain_tip") or "").lower() for h in hit):
                    matched.append(f)
            if matched:
                bits = [f"[{f.get('severity')}] {f.get('title')}" for f in matched[:8]]
                return (
                    "From the live threat scan: " + "; ".join(bits) + ". Fine if you trust it; Neutralize to contain the MAC.",
                    "threats",
                )
        if not threats:
            return (
                f"Threat hunt checked risky ports, new devices, gateway MAC, and probes — "
                f"{ts.get('findings', 0)} open finding(s) right now. No critical items waiting. "
                "Protect is not a full antivirus on every phone file; open Threats after each deep scan.",
                "threats",
            )
        bits = [f"[{f.get('severity')}] {f.get('title') or f.get('plain_tip')}" for f in threats[:8]]
        return (
            f"Open threats ({len(threats)}): " + "; ".join(bits) + ". Use Fine or Neutralize on the Threats tab.",
            "threats",
        )

    # Specific device
    dev = _find_device(ql, devices)
    if dev and any(
        w in ql
        for w in ("what is", "who is", "tell me", "about", "device", "this", "mikrotik", "router", "phone", "camera")
    ):
        name = dev.get("display_name") or dev.get("hostname") or dev.get("vendor") or dev.get("ip")
        return (
            f"{name}: type {dev.get('device_type') or 'unknown'}, vendor {dev.get('vendor') or '—'}, "
            f"IP {dev.get('ip') or '—'}, MAC {dev.get('mac') or '—'}, "
            f"ping {'timeout' if not dev.get('reachable') else str(dev.get('ping_ms')) + ' ms'}, "
            f"ports {dev.get('open_ports') or 'none seen'}, "
            f"traffic share {dev.get('traffic_share_pct') or 0}%. "
            f"Icon key: {dev.get('icon_key') or 'unknown'}.",
            "device",
        )

    if any(w in ql for w in ("bandwidth", "busy", "using", "traffic", "talker")) or (
        "who" in ql and any(w in ql for w in ("wifi", "wi-fi", "network", "using"))
    ):
        talkers = facts["talkers"]
        if not talkers:
            return (
                "I'm sampling traffic visible to this Protect box. No heavy talker in the last sample — check Traffic in a minute.",
                "bandwidth",
            )
        bits = [f"{t['name']} (~{t['pct']}%, {t['rate']} B/s)" for t in talkers]
        return (
            "Top observed talkers: " + "; ".join(bits) + ". Side-box sample — not a full switch mirror.",
            "bandwidth",
        )

    if any(w in ql for w in ("slow", "ping", "latency", "lag")):
        slow = [
            d
            for d in devices
            if d.get("reachable") and isinstance(d.get("ping_ms"), (int, float)) and float(d["ping_ms"]) >= 50
        ]
        slow.sort(key=lambda d: float(d.get("ping_ms") or 0), reverse=True)
        if not slow:
            return ("Latency looks healthy — answering devices are mostly under 50 ms.", "latency")
        return (
            "Higher latency: "
            + "; ".join(
                f"{d.get('display_name') or d.get('hostname') or d.get('ip')} {d.get('ping_ms')} ms" for d in slow[:6]
            ),
            "latency",
        )

    if any(w in ql for w in ("timeout", "offline", "not answering")):
        dead = [d for d in devices if d.get("ip") and not d.get("reachable")]
        if not dead:
            return ("No ping timeouts in the latest sweep.", "timeouts")
        return (
            "Not answering ping: "
            + ", ".join((d.get("display_name") or d.get("hostname") or d.get("ip") or d.get("mac")) for d in dead[:8]),
            "timeouts",
        )

    if any(w in ql for w in ("how many", "inventory", "list device", "devices on")) or ql.strip() in {
        "devices",
        "how many devices?",
        "how many devices",
    }:
        return (
            f"I see {len(devices)} devices on {sm.get('cidr') or 'the LAN'} "
            f"({pulse['reachable']} answering). Open Devices for icons, ping, ports, and vendor.",
            "inventory",
        )

    if any(w in ql for w in ("status", "score", "safe", "protected", "security status")):
        return (
            f"Security pulse {pulse['score']}/100 — {pulse['label']}. "
            f"Learning {pulse['learning_pct']:.0f}%. "
            f"{'Internet OK. ' if pulse['internet_ok'] else 'Internet looks down. '}"
            f"{pulse.get('threats', 0)} threat(s) open ({pulse.get('critical', 0)} critical). "
            "Open the Threats tab for Fine / Neutralize.",
            "status",
        )

    if any(w in ql for w in ("fine", "neutralize", "button", "block")):
        return (
            "Fine = you accept that finding for now. Neutralize = Protect contains that device MAC on the Protect path "
            "(side-box filter — for a full house block also remove it on your router/AP). I never auto-block.",
            "ui",
        )

    if any(w in ql for w in ("internet", "dns", "gateway", "wan", "router", "modem")):
        nh = facts["network_health"]
        bits = []
        for t in nh.get("targets") or []:
            bits.append(
                f"{t.get('host')} {t.get('avg_ms')} ms" if t.get("avg_ms") is not None else f"{t.get('host')} timeout"
            )
        gw_extra = ""
        try:
            import gateway_secure as gs

            gws = gs.public_status()
            if gws.get("configured"):
                last = gws.get("last_read") or {}
                tips = last.get("tips") or gws.get("tips_preview") or []
                tip_line = "; ".join((t.get("title") or "") for t in tips[:3] if isinstance(t, dict))
                gw_extra = (
                    f" Main gateway login is saved on this Protect box (read-only). "
                    f"Brand/IP from last read: {last.get('brand') or '—'} / {last.get('gateway_ip') or gws.get('gateway_ip') or '—'}. "
                    f"Atlas never changes the router unless fn.nology@gmail.com authorizes it."
                    + (f" Tips: {tip_line}." if tip_line else "")
                )
            else:
                gw_extra = (
                    " You can save the main gateway username/password under Insights → Main gateway "
                    "(read-only info for Atlas Secure tips)."
                )
        except Exception:
            pass
        return (
            f"From this Protect box: internet {'OK' if nh.get('internet_ok') else 'DOWN'}, "
            f"DNS {'OK' if nh.get('dns_ok') else 'failing'}, gateway {nh.get('gateway') or '—'}. "
            + ("RTT: " + "; ".join(bits) + ". " if bits else "")
            + "This is Protect’s path — not rewriting your router."
            + gw_extra,
            "internet",
        )

    if any(w in ql for w in ("what does", "what do you", "how it works", "can it", "can do", "actually do", "antivirus", "virus protection")):
        return (
            "Atlas Cyber Protect is a silent LAN side-box: sweeps devices, hunts threats (risky ports, new hosts, "
            "gateway MAC, probes), IDs devices with icons, samples bandwidth, checks internet/DNS, and asks you "
            "Fine/Neutralize. It is not your router and not a full antivirus agent on every PC file — "
            "a companion Protect Agent app on computers is the next layer for on-device virus scans. "
            + (f"Active: {'; '.join((facts.get('capabilities_active') or [])[:6])}." if facts.get("capabilities_active") else ""),
            "product",
        )

    if any(w in ql for w in ("service", "port", "camera", "ssh")):
        svcs = facts.get("services") or []
        if not svcs:
            return ("Port/service summary warms up after the deep LAN scan — check Traffic shortly.", "services")
        return (
            "Services on your LAN: " + "; ".join(f"{s.get('name')} ({s.get('count')} device(s))" for s in svcs[:8]),
            "services",
        )

    if any(w in ql for w in ("timeline", "new device", "appeared", "gone")):
        evs = st.get("device_events") or []
        if not evs:
            return ("Timeline fills when devices appear, go quiet, or return.", "timeline")
        return (
            "Recent device events: " + "; ".join((e.get("plain") or e.get("label") or e.get("kind")) for e in evs[:6]),
            "timeline",
        )

    # Open chat — still give a grounded stub; Atlas may polish
    return (
        f"Live on this Protect box: pulse {pulse['score']}/100 ({pulse['label']}), "
        f"{len(devices)} devices, {len(threats)} open threat(s). "
        f"Tabs: {', '.join(facts['ui_tabs'])}. Ask about a threat, a device name/IP, or a tab.",
        "open",
    )


def _local_ollama_polish(question: str, facts: dict, local_reply: str) -> str | None:
    """Advanced local-LLM rewrite (home / GPU boxes only).

    When acp.env sets ACP_LOCAL_LLM=ollama://<model> the guardian talks
    directly to the box's own Ollama (e.g. qwen2.5:14b on the RX 7800 XT)
    instead of the funnel's CPU model — same grounded facts, much smarter
    answers. Returns None when not configured or unreachable, so the funnel
    path stays the fallback for ordinary boxes.
    """
    spec = (os.environ.get("ACP_LOCAL_LLM") or "").strip()
    if not spec.startswith("ollama://"):
        return None
    model = spec.split("ollama://", 1)[1].strip() or ""
    if not model:
        return None
    host = (os.environ.get("ACP_LOCAL_OLLAMA_HOST") or "http://127.0.0.1:11434").rstrip("/")
    system = (
        "You are Atlas Cyber Protect's guardian for a paying customer. "
        "CRITICAL: Use ONLY the snapshot facts. Never invent devices, IPs, threats, or scores. "
        "If snapshot.local_truth is present, keep the same facts — you may only rephrase clearer. "
        "Never mention Factory, prices, ZAR, or becoming the gateway. "
        "Never claim you auto-blocked a device — customer chooses Fine or Neutralize. "
        "UI tabs are Overview, Devices, Traffic, Insights, Threats, Ask guardian, How it works. "
        "Keep answers under 140 words."
    )
    devices = facts.get("devices") if isinstance(facts.get("devices"), list) else []
    threats = facts.get("threats") if isinstance(facts.get("threats"), list) else []
    user = (
        f"Customer question: {question}\n"
        f"MUST KEEP THESE FACTS (local_truth): {local_reply[:900]}\n"
        f"Pulse: {json.dumps(facts.get('pulse') or {})[:400]}\n"
        f"Threat summary: {json.dumps(facts.get('threat_summary') or {})[:400]}\n"
        f"Threats: {json.dumps(threats[:10])[:1000]}\n"
        f"Devices (sample): {json.dumps(devices[:12])[:1500]}\n"
        f"Talkers: {json.dumps(facts.get('talkers') or [])[:400]}\n"
        f"Network health: {json.dumps(facts.get('network_health') or {})[:400]}\n"
    )
    try:
        payload = {
            "model": model,
            "messages": [
                {"role": "system", "content": system},
                {"role": "user", "content": user},
            ],
            "stream": False,
            "options": {"num_predict": 280, "temperature": 0.4},
        }
        req = urllib.request.Request(
            f"{host}/api/chat",
            data=json.dumps(payload).encode(),
            headers={"Content-Type": "application/json"},
            method="POST",
        )
        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:
            out = json.loads(resp.read().decode("utf-8", errors="replace") or "{}")
        text = ((out.get("message") or {}).get("content") or "").strip()
        return text or None
    except Exception:
        return None


def _atlas_polish(question: str, facts: dict, local_reply: str, enroll_token: str) -> str | None:
    """Optional LLM rewrite — MUST stay faithful to facts + local_reply."""
    if not enroll_token:
        return None
    payload = {
        "enroll_token": enroll_token,
        "message": question,
        "snapshot": {
            "pulse": facts.get("pulse"),
            "scan_meta": facts.get("scan_meta"),
            "devices": facts.get("devices"),
            "findings": [t.get("plain") or t.get("title") for t in (facts.get("threats") or [])],
            "threats": facts.get("threats"),
            "threat_summary": facts.get("threat_summary"),
            "network_health": facts.get("network_health"),
            "talkers": facts.get("talkers"),
            "services": facts.get("services"),
            "capabilities_active": facts.get("capabilities_active"),
            "ui_tabs": facts.get("ui_tabs"),
            "local_truth": local_reply,
        },
    }
    try:
        data = json.dumps(payload).encode()
        req = urllib.request.Request(
            f"{atlas_base()}/api/cyber-protect/assistant",
            data=data,
            headers={"Content-Type": "application/json"},
            method="POST",
        )
        ctx = ssl.create_default_context()
        ctx.check_hostname = False
        ctx.verify_mode = ssl.CERT_NONE
        with urllib.request.urlopen(req, timeout=35, context=ctx) as resp:
            out = json.loads(resp.read().decode("utf-8", errors="replace") or "{}")
        text = (out.get("reply") or out.get("text") or "").strip()
        return text or None
    except Exception:
        return None


def ask(question: str, *, enroll_token: str = "") -> dict[str, Any]:
    st = load_state()
    live_feed(st)
    local, intent = _local_answer(question, st)
    facts = _facts_pack(st)
    source = "local"
    reply = local
    # Only polish open-ended chat with Atlas — never replace factual intents
    if intent == "open" and (question or "").strip():
        # Advanced boxes (home GPU server etc.): prefer the box's own big local
        # model (ACP_LOCAL_LLM=ollama://qwen2.5:14b) — smarter, still grounded.
        remote = _local_ollama_polish(question, facts, local) or _atlas_polish(question, facts, local, enroll_token)
        if remote:
            # Guardrail: if model invents empty network, keep local
            if len(facts.get("devices") or []) > 0 and "0 devices" in remote.lower() and "16" not in remote:
                reply = local
            else:
                reply = remote
                source = "atlas"
    feed = list(st.get("guardian_feed") or [])
    feed.insert(0, {"ts": time.time(), "text": f"You asked: {(question or 'status')[:120]}", "kind": "user"})
    feed.insert(0, {"ts": time.time(), "text": reply, "kind": "guardian"})
    st["guardian_feed"] = feed[:60]
    save_state(st)
    return {
        "ok": True,
        "reply": reply,
        "source": source,
        "intent": intent,
        "pulse": st.get("guardian_pulse") or security_score(st),
    }


def refresh_living_state(st: dict | None = None) -> dict[str, Any]:
    st = st if st is not None else load_state()
    live_feed(st)
    save_state(st)
    return st
