from __future__ import annotations

import json
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from queue import Queue

from agent_guard.events import Event
from agent_guard.notifier import journal_send


def _reason(alert: dict) -> str:
    annotations = alert.get("annotations") or {}
    labels = alert.get("labels") or {}
    summary = str(annotations.get("summary") or "").strip()
    if summary:
        return summary
    pieces = [str(labels.get("alertname") or "").strip(), str(alert.get("valueString") or "").strip()]
    return " ".join(piece for piece in pieces if piece).strip() or "Grafana alert firing"


DEFAULT_GRAFANA_URL = "http://127.0.0.1:3000"
OVERVIEW_DASHBOARD_PATH = "/d/system-monitor-overview"


def _info_url(alert: dict, base: str) -> str:
    # Grafana-managed rules often omit generatorURL, and panelURL/dashboardURL
    # appear only with both dashboard+panel annotations; always fall back to the
    # overview dashboard so the Info button is never dead.
    for key in ("panelURL", "dashboardURL", "generatorURL"):
        value = str(alert.get(key) or "").strip()
        if value:
            return value
    return base + OVERVIEW_DASHBOARD_PATH


def _remediation(alert: dict) -> str:
    labels = alert.get("labels") or {}
    remediation = str(labels.get("remediation") or "none").strip().lower()
    return remediation if remediation in {"kill", "clean", "none"} else "none"


def map_alert(alert: dict, base: str = DEFAULT_GRAFANA_URL) -> Event | None:
    if not isinstance(alert, dict):
        raise TypeError("alert must be a dict")
    status = str(alert.get("status") or "").strip().lower()
    if status == "resolved":
        return None
    if status != "firing":
        raise ValueError("unsupported alert status")
    labels = alert.get("labels") or {}
    severity = str(labels.get("severity") or "").strip().lower()
    return Event(
        "critical" if severity == "critical" else "warning",
        "grafana",
        _reason(alert),
        info_url=_info_url(alert, base),
        remediation=_remediation(alert),
    )


def _parse_events(body: bytes) -> list[Event]:
    payload = json.loads(body.decode("utf-8"))
    alerts = payload.get("alerts")
    if not isinstance(alerts, list):
        raise TypeError("alerts must be a list")
    base = str(payload.get("externalURL") or DEFAULT_GRAFANA_URL).strip().rstrip("/") or DEFAULT_GRAFANA_URL
    events = []
    for alert in alerts:
        event = map_alert(alert, base)
        if event is not None:
            events.append(event)
    return events


def handle_payload(queue: Queue, body: bytes) -> int:
    try:
        events = _parse_events(body)
    except (UnicodeDecodeError, json.JSONDecodeError, TypeError, ValueError):
        journal_send(SM_TIER="warning", SM_SOURCE="grafana", SM_ACTION="bad_webhook", SM_TARGET="127.0.0.1:9099")
        return 400
    for event in events:
        queue.put(event)
    return 202


def make_server(queue: Queue, host: str = "127.0.0.1", port: int = 9099) -> ThreadingHTTPServer:
    if host != "127.0.0.1":
        raise ValueError("agent-guard webhook must bind 127.0.0.1 only")

    class WebhookHandler(BaseHTTPRequestHandler):
        def do_POST(self) -> None:
            if self.path != "/":
                self.send_error(404)
                return
            length = int(self.headers.get("Content-Length") or "0")
            body = self.rfile.read(length)
            status = handle_payload(queue, body)
            if status == 400:
                self.send_response(400)
                self.end_headers()
                self.wfile.write(b"bad webhook\n")
                return
            self.send_response(status)
            self.end_headers()
            self.wfile.write(b"accepted\n")

        def log_message(self, _format: str, *_args) -> None:
            return

    return ThreadingHTTPServer((host, port), WebhookHandler)


def follow(queue: Queue, stop, host: str = "127.0.0.1", port: int = 9099) -> None:
    server = make_server(queue, host=host, port=port)
    server.timeout = 0.5
    with server:
        while not stop.is_set():
            server.handle_request()
