from __future__ import annotations

import json
import os
import socket
import subprocess
import threading
import time
from collections import deque
from queue import Empty, Queue

from agent_guard.config import load_config
from agent_guard.culprit import rescan_culprit, terminate_target
from agent_guard.events import Event
from agent_guard.events import fswatch, journald, ports, webhook
from agent_guard.notifier import action_monitor, clean_tmp_safe, journal_send, notify_event
from agent_guard.state import State


def sd_notify(message: str) -> None:
    path = os.environ.get("NOTIFY_SOCKET")
    if not path:
        return
    addr = ("\0" + path[1:]) if path.startswith("@") else path
    with socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) as sock:
        sock.sendto(message.encode(), addr)


POLL_SNAPSHOT_MAX_EVENTS = 200


class PollSnapshot:
    """Thread-safe view of the last POLL_SNAPSHOT_MAX_EVENTS events this daemon handled."""

    def __init__(self, max_events: int = POLL_SNAPSHOT_MAX_EVENTS) -> None:
        self._events: deque[dict] = deque(maxlen=max_events)
        self._lock = threading.Lock()

    def record(self, event: Event) -> None:
        entry = {
            "tier": event.tier,
            "source": event.source,
            "reason": event.reason,
            "remediation": event.remediation,
        }
        if event.culprit is not None:
            entry["culprit"] = event.culprit
        with self._lock:
            self._events.append(entry)

    def response(self) -> dict:
        with self._lock:
            events = list(self._events)
        return {
            "events": events,
            "culprits": [event["culprit"] for event in events if "culprit" in event],
        }


def handle_ipc(conn, queue: Queue, snapshot: PollSnapshot) -> None:
    msg = json.loads(conn.recv(65536).decode().strip())
    if not isinstance(msg, dict):
        raise TypeError("ipc message must be a JSON object")
    if msg.get("type") == "poll":
        conn.sendall(json.dumps(snapshot.response()).encode() + b"\n")
        return
    ev = msg["event"]
    actions = ev.get("actions") or []
    remediation = ev.get("remediation") or remediation_from_actions(actions)
    queue.put(Event(
        ev["tier"],
        ev["source"],
        ev["reason"],
        ev.get("culprit"),
        actions,
        ev.get("info_url"),
        remediation,
    ))


def socket_server(path, queue: Queue, stop, snapshot: PollSnapshot) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    try:
        path.unlink()
    except FileNotFoundError:
        pass
    with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as srv:
        srv.bind(str(path))
        os.chmod(path, 0o666)
        srv.listen(10)
        srv.settimeout(1)
        while not stop.is_set():
            try:
                conn, _ = srv.accept()
            except TimeoutError:
                continue
            try:
                with conn:
                    handle_ipc(conn, queue, snapshot)
            except (json.JSONDecodeError, KeyError, TypeError, UnicodeDecodeError):
                journal_send(SM_TIER="warning", SM_SOURCE="agent-guard", SM_ACTION="bad_ipc", SM_TARGET="socket")
            except OSError:
                journal_send(SM_TIER="warning", SM_SOURCE="agent-guard", SM_ACTION="ipc_transport", SM_TARGET="socket")


def remediation_from_actions(actions: list[str]) -> str:
    normalized = " ".join(action.lower() for action in actions)
    if "kill" in normalized:
        return "kill"
    if "clean" in normalized:
        return "clean"
    return "none"


def open_info_url(url: str) -> None:
    for cmd in (["xdg-open", url], ["gio", "open", url]):
        try:
            subprocess.run(cmd, check=False, timeout=5)
            return
        except OSError:
            continue


def handle_event(event: Event, cfg) -> None:
    info = ["Info"] if event.info_url else []  # no dead Info button when there is no URL
    if event.remediation == "kill":
        culprit = rescan_culprit(cfg.protect)
        if culprit:
            event.culprit = culprit
            event.reason = f"{event.reason}; fastest grower: {culprit['name']} ({len(culprit.get('pids', []))} procs)"
            event.actions = info + [f"Kill {culprit['name']} ({len(culprit.get('pids', []))} procs)"]
        else:
            event.actions = info + ["Dismiss"]
    elif event.remediation == "clean":
        event.actions = info + ["Clean /tmp"]
    else:
        event.actions = info + (event.actions or ["Dismiss"])
    notify_event(event)


def handle_action(event: Event, action: str, cfg) -> None:
    if action.startswith("info"):
        if event.info_url:
            open_info_url(event.info_url)
            notify_event(event)
        return
    if action.startswith("kill") and event.culprit:
        survivors = terminate_target(event.culprit, cfg.protect)
        journal_send(SM_TIER=event.tier, SM_SOURCE=event.source, SM_ACTION="kill", SM_TARGET=event.culprit.get("name", "unknown"), SM_SURVIVORS=str(len(survivors)))
    elif action.startswith("clean"):
        removed = clean_tmp_safe(uid=1000, older_than_seconds=cfg.tmp_stale_hours * 3600)
        journal_send(SM_TIER=event.tier, SM_SOURCE=event.source, SM_ACTION="clean_tmp", SM_TARGET="/tmp", SM_REMOVED=str(removed))


def main() -> int:
    cfg = load_config()
    state = State(cfg.state_path)
    queue: Queue = Queue()
    stop = threading.Event()
    snapshot = PollSnapshot()
    threads = [
        threading.Thread(target=journald.follow, args=(queue, stop), daemon=True),
        threading.Thread(target=fswatch.follow, args=(queue, cfg.watched_paths, stop), daemon=True),
        threading.Thread(target=ports.tick, args=(queue, state, cfg, stop), daemon=True),
        threading.Thread(target=webhook.follow, args=(queue, stop), daemon=True),
        threading.Thread(target=socket_server, args=(cfg.socket_path, queue, stop, snapshot), daemon=True),
        threading.Thread(target=action_monitor, args=(lambda event, action: handle_action(event, action, cfg), stop), daemon=True),
    ]
    for thread in threads:
        thread.start()
    sd_notify("READY=1")
    last_watchdog = 0.0
    try:
        while True:
            now = time.time()
            if now - last_watchdog > 10:
                sd_notify("WATCHDOG=1")
                last_watchdog = now
            try:
                event = queue.get(timeout=1)
            except Empty:
                continue
            cooldown = cfg.cooldowns.get(event.source, cfg.cooldowns.get("default", 300))
            key = f"{event.source}:{event.reason}"
            if state.allow_now(key, cooldown):
                handle_event(event, cfg)
                snapshot.record(event)
    except KeyboardInterrupt:
        stop.set()
        return 0


if __name__ == "__main__":
    raise SystemExit(main())
