#!/usr/bin/env python3
"""cluster-watchdog: notice a NotReady node or a sick Longhorn volume and tell a human.

Runs as a Kubernetes CronJob every 5 minutes. Standard library only (urllib + hmac/hashlib),
so it runs straight from the stock python:3.12-slim image with this file mounted from a
ConfigMap -- no custom image, no pip.

What it checks
  * Every node's Ready condition. A node that is not Ready for longer than NODE_GRACE_MINUTES
    is a problem (a quick reboot is not).
  * Every Longhorn volume. `faulted` is a problem immediately; `degraded` while attached is a
    problem once it has lasted LONGHORN_GRACE_MINUTES (a replica rebuild after a node returns
    is normal for a while). Detached volumes with robustness `unknown` are normal and ignored.

How it alerts (dedup so a 4-week outage is not 8000 messages)
  * State lives in the ConfigMap STATE_CONFIGMAP in the job's own namespace: one entry per
    problem key with first_seen / last_notified.
  * A message goes out when a problem is new, when one has recovered, or as a reminder when a
    problem has persisted REMIND_HOURS since the last message. One run = at most one message.
  * Channels are Azure Communication Services (ACS) email and/or SMS, signed with the same
    access-key HMAC-SHA256 scheme as the media SMS notifier
    (learn.microsoft.com/azure/communication-services/tutorials/hmac-header-tutorial).
    With no channel configured it logs what it would have sent and still tracks state, and it
    notices when a channel appears later and re-sends every open problem once.

Environment (all optional except the namespace)
  STATE_NAMESPACE          namespace holding the state ConfigMap (downward API)
  STATE_CONFIGMAP          default cluster-watchdog-state
  NODE_GRACE_MINUTES       default 10
  LONGHORN_GRACE_MINUTES   default 60
  REMIND_HOURS             default 24
  ACS_ENDPOINT             https://<resource>.communication.azure.com
  ACS_ACCESS_KEY           the resource's access key (base64 as shown in the portal)
  ACS_EMAIL_FROM           a verified sender address on the resource's email domain
  ACS_EMAIL_TO             comma-separated recipient addresses
  ACS_SMS_FROM             an SMS-capable number on the resource (E.164)
  ACS_SMS_TO               comma-separated recipient numbers (E.164)
"""

import base64
import hashlib
import hmac
import json
import os
import re
import ssl
import sys
from datetime import datetime, timedelta, timezone
from urllib import error as urlerror
from urllib import request as urlrequest

SA_DIR = "/var/run/secrets/kubernetes.io/serviceaccount"
API = f"https://{os.environ.get('KUBERNETES_SERVICE_HOST', 'kubernetes.default.svc')}:{os.environ.get('KUBERNETES_SERVICE_PORT', '443')}"

STATE_NAMESPACE = os.environ.get("STATE_NAMESPACE") or "cluster-watchdog"
STATE_CONFIGMAP = os.environ.get("STATE_CONFIGMAP") or "cluster-watchdog-state"
def _env_float(name, default):
    # A bad tuning value (e.g. NODE_GRACE_MINUTES="10m") must not raise at import, before
    # main()'s try/except exists, and crash-loop the CronJob into silence. Fall back to the
    # default and log instead.
    raw = os.environ.get(name)
    if not raw:
        return float(default)
    try:
        return float(raw)
    except (ValueError, TypeError):
        print(f"cluster-watchdog: {name}={raw!r} is not a number; using default {default}", flush=True)
        return float(default)


NODE_GRACE = timedelta(minutes=_env_float("NODE_GRACE_MINUTES", 10))
LONGHORN_GRACE = timedelta(minutes=_env_float("LONGHORN_GRACE_MINUTES", 60))
REMIND = timedelta(hours=_env_float("REMIND_HOURS", 24))

ACS_ENDPOINT = (os.environ.get("ACS_ENDPOINT") or "").rstrip("/")
ACS_ACCESS_KEY = os.environ.get("ACS_ACCESS_KEY") or ""
ACS_EMAIL_FROM = os.environ.get("ACS_EMAIL_FROM") or ""
ACS_SMS_FROM = os.environ.get("ACS_SMS_FROM") or ""


def _split(value):
    return [v for v in re.split(r"[,;\s]+", value or "") if v]


EMAIL_TO = _split(os.environ.get("ACS_EMAIL_TO"))
SMS_TO = _split(os.environ.get("ACS_SMS_TO"))
EMAIL_API_VERSION = "2023-03-31"
SMS_API_VERSION = "2021-03-07"

_DAYS = ["Mon", "Tue", "Wed", "Thu", "Fri", "Sat", "Sun"]
_MONTHS = ["Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct", "Nov", "Dec"]


def log(message):
    print(f"{datetime.now(timezone.utc).strftime('%Y-%m-%dT%H:%M:%SZ')} {message}", flush=True)


def parse_time(value):
    """Kubernetes RFC 3339 timestamp -> aware datetime; None when absent/blank."""
    if not value:
        return None
    return datetime.strptime(value[:19], "%Y-%m-%dT%H:%M:%S").replace(tzinfo=timezone.utc)


def fmt(dt):
    return dt.strftime("%Y-%m-%d %H:%MZ")


# --- Kubernetes API (in-cluster service account) -------------------------------------------

_SSL = ssl.create_default_context(cafile=f"{SA_DIR}/ca.crt")
with open(f"{SA_DIR}/token", encoding="utf-8") as _f:
    _TOKEN = _f.read().strip()


def k8s(method, path, body=None):
    data = json.dumps(body).encode("utf-8") if body is not None else None
    headers = {"Authorization": f"Bearer {_TOKEN}", "Accept": "application/json"}
    if data is not None:
        headers["Content-Type"] = "application/json"
    req = urlrequest.Request(f"{API}{path}", data=data, headers=headers, method=method)
    with urlrequest.urlopen(req, timeout=20, context=_SSL) as resp:
        return json.load(resp)


# --- Checks ---------------------------------------------------------------------------------

def check_nodes(now):
    """Return ({problem_key: {"label", "text"}}, summary_line)."""
    problems, ready, total = {}, 0, 0
    for node in k8s("GET", "/api/v1/nodes").get("items", []):
        total += 1
        name = node["metadata"]["name"]
        cond = next((c for c in node.get("status", {}).get("conditions", []) if c.get("type") == "Ready"), None)
        status = cond.get("status") if cond else "Unknown"
        if status == "True":
            ready += 1
            continue
        since = parse_time(cond.get("lastTransitionTime")) if cond else None
        if since is None or now - since >= NODE_GRACE:
            reason = (cond or {}).get("reason") or "no Ready condition"
            since_text = fmt(since) if since else "unknown"
            problems[f"node/{name}"] = {"label": name, "text": f"node {name} is NotReady since {since_text} ({reason})"}
    return problems, f"nodes {ready}/{total} Ready"


def check_longhorn(now):
    """Return ({problem_key: {"label", "text"}}, summary_line)."""
    problems, counts = {}, {}
    volumes = k8s("GET", "/apis/longhorn.io/v1beta2/namespaces/longhorn-system/volumes").get("items", [])
    for vol in volumes:
        status = vol.get("status", {})
        robustness = status.get("robustness") or "unknown"
        state = status.get("state") or "unknown"
        counts[robustness] = counts.get(robustness, 0) + 1
        kube = status.get("kubernetesStatus") or {}
        label = f"{kube.get('namespace') or '?'}/{kube.get('pvcName') or vol['metadata']['name']}"
        key = f"volume/{vol['metadata']['name']}"
        if robustness == "faulted":
            problems[key] = {"label": label, "text": f"Longhorn volume {label} is FAULTED (all replicas lost)"}
        elif robustness == "degraded" and state == "attached":
            # Only alert once it has been degraded past the grace period. A missing
            # lastDegradedAt means Longhorn just started the rebuild, so wait rather than
            # page during the exact routine rebuild the grace exists to suppress (faulted,
            # above, has no grace and still alerts immediately).
            since = parse_time(status.get("lastDegradedAt"))
            if since is not None and now - since >= LONGHORN_GRACE:
                problems[key] = {"label": label, "text": f"Longhorn volume {label} is degraded since {fmt(since)} (replica missing; a rebuild may be in progress)"}
    summary = ", ".join(f"{n} {r}" for r, n in sorted(counts.items()))
    return problems, f"Longhorn volumes: {summary or 'none'}"


# --- State ConfigMap ------------------------------------------------------------------------

def load_state():
    path = f"/api/v1/namespaces/{STATE_NAMESPACE}/configmaps/{STATE_CONFIGMAP}"
    try:
        raw = (k8s("GET", path).get("data") or {}).get("state.json") or "{}"
        data = json.loads(raw)
        # A hand-edited ConfigMap could hold valid JSON of the wrong shape (a list, a string,
        # null). Only a dict is usable; anything else starts from empty rather than crashing.
        return data if isinstance(data, dict) else {}
    except urlerror.HTTPError as exc:
        if exc.code == 404:
            return {}
        raise
    except (ValueError, TypeError, AttributeError) as exc:
        # A hand-edited or corrupt state ConfigMap must not crash-loop the watchdog into
        # silence; start from empty state (a fresh problem just re-alerts) and carry on.
        log(f"load_state: unreadable state ConfigMap ({exc}); starting from empty")
        return {}


def save_state(state):
    body = {
        "apiVersion": "v1",
        "kind": "ConfigMap",
        "metadata": {"name": STATE_CONFIGMAP, "namespace": STATE_NAMESPACE},
        "data": {"state.json": json.dumps(state, indent=1, sort_keys=True)},
    }
    base = f"/api/v1/namespaces/{STATE_NAMESPACE}/configmaps"
    try:
        k8s("PUT", f"{base}/{STATE_CONFIGMAP}", body)
    except urlerror.HTTPError as exc:
        if exc.code != 404:
            raise
        k8s("POST", base, body)


# --- Azure Communication Services -----------------------------------------------------------

def _rfc1123(dt):
    # Locale-independent RFC 1123 GMT timestamp for the x-ms-date header.
    return (f"{_DAYS[dt.weekday()]}, {dt.day:02d} {_MONTHS[dt.month - 1]} {dt.year:04d} "
            f"{dt.hour:02d}:{dt.minute:02d}:{dt.second:02d} GMT")


def acs_post(path, payload):
    """POST a JSON body to ACS with the access-key HMAC-SHA256 signature. Returns True on 2xx."""
    url = f"{ACS_ENDPOINT}{path}"
    host = url.split("://", 1)[1].split("/", 1)[0]
    body = json.dumps(payload).encode("utf-8")
    content_hash = base64.b64encode(hashlib.sha256(body).digest()).decode()
    date = _rfc1123(datetime.now(timezone.utc))
    string_to_sign = f"POST\n{path}\n{date};{host};{content_hash}"
    signature = base64.b64encode(
        hmac.new(base64.b64decode(ACS_ACCESS_KEY), string_to_sign.encode("utf-8"), hashlib.sha256).digest()
    ).decode()
    headers = {
        "x-ms-date": date,
        "x-ms-content-sha256": content_hash,
        "Authorization": f"HMAC-SHA256 SignedHeaders=x-ms-date;host;x-ms-content-sha256&Signature={signature}",
        "Content-Type": "application/json",
    }
    req = urlrequest.Request(url, data=body, headers=headers, method="POST")
    try:
        with urlrequest.urlopen(req, timeout=20) as resp:
            log(f"acs: {path.split('?')[0]} accepted ({resp.status})")
            return True
    except urlerror.HTTPError as exc:
        detail = exc.read().decode("utf-8", "replace")[:200]
        log(f"acs: {path.split('?')[0]} FAILED (HTTP {exc.code}: {detail})")
    except Exception as exc:  # network/timeout -- report and let the next run retry
        log(f"acs: {path.split('?')[0]} FAILED ({exc})")
    return False


def channels():
    """Which delivery channels are fully configured right now."""
    have = []
    if ACS_ENDPOINT and ACS_ACCESS_KEY:
        if ACS_EMAIL_FROM and EMAIL_TO:
            have.append("email")
        if ACS_SMS_FROM and SMS_TO:
            have.append("sms")
    return have


def send(subject, text):
    """Deliver on every configured channel. Returns True when at least one delivery succeeded
    (or when nothing is configured, so state still advances in log-only mode)."""
    have = channels()
    if not have:
        log(f"notify (no channel configured -- create the acs-notify secret): {subject}\n{text}")
        return True
    ok = False
    if "email" in have:
        ok |= acs_post(f"/emails:send?api-version={EMAIL_API_VERSION}", {
            "senderAddress": ACS_EMAIL_FROM,
            "recipients": {"to": [{"address": a} for a in EMAIL_TO]},
            "content": {"subject": subject, "plainText": text},
            "userEngagementTrackingDisabled": True,
        })
    if "sms" in have:
        ok |= acs_post(f"/sms?api-version={SMS_API_VERSION}", {
            "from": ACS_SMS_FROM,
            "smsRecipients": [{"to": n} for n in SMS_TO],
            "message": f"{subject}\n{text}"[:300],
            "smsSendOptions": {"enableDeliveryReport": False},
        })
    return ok


# --- Main -----------------------------------------------------------------------------------

def main():
    now = datetime.now(timezone.utc)
    current, summaries = {}, []
    for check in (check_nodes, check_longhorn):
        try:
            found, summary = check(now)
        except Exception as exc:  # one broken check must not hide the other -- and a check
            # that fails every run is itself an alertable problem (the watchdog is blind to
            # that dimension), not something to swallow into the logs.
            log(f"{check.__name__}: FAILED ({exc})")
            key = f"check/{check.__name__}"
            found = {key: {"label": check.__name__, "text": f"watchdog {check.__name__} is failing: {exc}"}}
            summary = f"{check.__name__} FAILED: {exc}"
        current.update(found)
        summaries.append(summary)
    summary = "; ".join(summaries)

    state = load_state()
    # Ignore any non-dict entry (a corrupt/hand-edited value) rather than crash on .get() later.
    tracked = {k: v for k, v in state.items() if not k.startswith("_") and isinstance(v, dict)}
    have = channels()
    channels_changed = state.get("_channels") != have

    new = [k for k in current if k not in tracked]
    gone = [k for k in tracked if k not in current]
    due = []
    for key in current:
        if key in new:
            continue
        last = parse_time(tracked[key].get("last_notified"))
        if channels_changed or last is None or now - last >= REMIND:
            due.append(key)

    log(f"{summary}; problems={len(current)} new={len(new)} reminders={len(due)} recovered={len(gone)}")

    # A brand-new problem's first_seen is shown in reminders; use the current run's time but
    # only PERSIST it once the message is actually delivered (below), so a failed first send
    # leaves the key un-tracked and it is re-announced as NEW next run, not as a STILL reminder.
    first_seen_now = now.strftime("%Y-%m-%dT%H:%M:%SZ")

    delivered = True  # nothing to send counts as success (state may still advance channels)
    if new or due or gone:
        lines = []
        for key in new:
            lines.append(f"NEW: {current[key]['text']}")
        for key in due:
            prior = tracked[key].get("first_seen", first_seen_now)
            lines.append(f"STILL: {current[key]['text']} (first seen {prior})")
        for key in gone:
            lines.append(f"RECOVERED: {tracked[key].get('label', key)} (was: {tracked[key].get('text', '?')})")
        lines.append("")
        lines.append(f"Cluster now: {summary}.")
        if current:
            subject = f"[homelab] {len(current)} problem(s): {', '.join(sorted(v['label'] for v in current.values()))}"
        else:
            subject = "[homelab] all clear"
        delivered = send(subject, "\n".join(lines))
        if delivered:
            for key in new:
                tracked[key] = {"first_seen": first_seen_now, **current[key]}
            for key in new + due:
                tracked[key]["last_notified"] = first_seen_now
            for key in gone:
                tracked.pop(key, None)

    # Record the delivery channel set only after a successful (or unneeded) send; otherwise keep
    # the previous value so channels_changed stays true next run and the resend is retried, not lost.
    tracked["_channels"] = have if delivered else state.get("_channels")
    save_state(tracked)


if __name__ == "__main__":
    try:
        main()
    except Exception as exc:  # a crash here is itself worth seeing in `kubectl logs`
        log(f"cluster-watchdog: FAILED ({exc})")
        sys.exit(1)
