#!/usr/bin/python3
# Nagios/Icinga plugin to check the age of the newest data an InfluxDB 2 query returns.
#
# Copyright (C) Thomas Wagner <wagner-thomas@gmx.at>
# SPDX-License-Identifier: GPL-2.0-or-later

import argparse
import csv
import datetime
import io
import json
import os
import re
import socket
import ssl
import sys
import urllib.error
import urllib.parse
import urllib.request

OK = 0
WARNING = 1
CRITICAL = 2
UNKNOWN = 3

STATE_TEXT = {OK: "OK", WARNING: "WARNING", CRITICAL: "CRITICAL", UNKNOWN: "UNKNOWN"}
STATE_BY_NAME = {"ok": OK, "warning": WARNING, "critical": CRITICAL, "unknown": UNKNOWN}

VERSION = "0.1"
DEFAULT_PORT = 8086
DEFAULT_TIMEOUT = 10
DEFAULT_WARNING = "10m"
DEFAULT_CRITICAL = "30m"

DURATION_UNITS = {"": 1, "s": 1, "m": 60, "h": 3600, "d": 86400}

# RFC 3339 as InfluxDB writes it, with up to nine digits of the second's fraction
TIMESTAMP_RE = re.compile(r"^(\d{4}-\d\d-\d\d)[Tt ](\d\d:\d\d:\d\d)(?:\.(\d+))?([Zz]|[+-]\d\d:\d\d)$")


class PluginError(Exception):
    def __init__(self, state, message):
        super().__init__(message)
        self.state = state
        self.message = message


def parse_number(text):
    """A threshold bound in seconds; a unit s, m, h or d may follow the number."""
    match = re.fullmatch(r"([-+]?\d+(?:\.\d*)?|[-+]?\.\d+)([smhd]?)", text.strip())
    if not match:
        raise ValueError(text)
    return float(match.group(1)) * DURATION_UNITS[match.group(2)]


class Range:
    """Threshold range as defined by the monitoring plugins guidelines, in seconds."""

    def __init__(self, start, end, inside):
        self.start = start
        self.end = end
        self.inside = inside
        self.spec = self._normalized()

    def _normalized(self):
        """Threshold as perfdata carries it: graphers want a plain number, so one-sided ranges become their bound."""
        if not self.inside and self.start == float("-inf") and self.end == float("inf"):
            return ""
        if not self.inside and self.end == float("inf"):
            return format_number(self.start)
        if not self.inside and self.start in (0, float("-inf")) and self.end != float("inf"):
            return format_number(self.end)
        low = "~" if self.start == float("-inf") else format_number(self.start)
        high = "" if self.end == float("inf") else format_number(self.end)
        return ("@" if self.inside else "") + "%s:%s" % (low, high)

    @classmethod
    def parse(cls, spec):
        raw = spec
        inside = spec.startswith("@")
        if inside:
            spec = spec[1:]
        if ":" in spec:
            low, high = spec.split(":", 1)
        else:
            low, high = "0", spec
        try:
            start = float("-inf") if low in ("~", "") else parse_number(low)
            end = float("inf") if high == "" else parse_number(high)
        except ValueError:
            hint = ""
            if raw.startswith("/") and ":" in raw:
                # bash turns an unquoted ~:10m into /home/user:10m
                hint = " (did the shell expand an unquoted ~? Quote the range: '~:...')"
            raise PluginError(UNKNOWN, "invalid threshold %r%s" % (raw, hint))
        if start > end:
            raise PluginError(UNKNOWN, "invalid threshold %r: start is above end" % raw)
        return cls(start, end, inside)

    def breached(self, value):
        within = self.start <= value <= self.end
        return within if self.inside else not within


def format_number(value):
    if float(value).is_integer():
        return "%d" % value
    return ("%.3f" % value).rstrip("0").rstrip(".")


def format_duration(seconds):
    """Human readable duration with its two most significant units, e.g. 2d 3h or 5m 10s."""
    sign = "-" if seconds < 0 else ""
    rest = int(round(abs(seconds)))
    parts = []
    for unit, size in (("d", 86400), ("h", 3600), ("m", 60), ("s", 1)):
        if rest >= size or (unit == "s" and not parts):
            parts.append("%d%s" % (rest // size, unit))
            rest %= size
    return sign + " ".join(parts[:2])


def perfdata(label, value, uom="", warn=None, crit=None, minimum=None):
    fields = [warn.spec if warn else "", crit.spec if crit else "",
              "" if minimum is None else format_number(minimum)]
    while fields and fields[-1] == "":
        fields.pop()
    out = "%s=%s%s" % (label, format_number(value), uom)
    if fields:
        out += ";" + ";".join(fields)
    return out


def parse_timestamp(text):
    """An RFC 3339 timestamp of InfluxDB, e.g. 2026-10-02T07:10:45.123456789Z, as an aware datetime."""
    match = TIMESTAMP_RE.match(text.strip())
    if not match:
        raise ValueError(text)
    date, clock, fraction, zone = match.groups()
    # datetime takes exactly six digits of the fraction before Python 3.11
    fraction = (fraction or "")[:6].ljust(6, "0")
    zone = "+00:00" if zone in ("Z", "z") else zone
    return datetime.datetime.fromisoformat("%sT%s.%s%s" % (date, clock, fraction, zone))


def newest_timestamp(text):
    """The newest _time of a CSV answer of the Flux query API, or None for an empty answer.

    Every table of the answer starts with its own header line, and tables may
    have different columns, so the position of _time is taken from each header.
    Blank lines separate the tables, and #annotation lines are skipped.
    """
    newest = None
    column = None
    tables = 0
    tables_with_time = 0
    for row in csv.reader(io.StringIO(text)):
        if not any(cell.strip() for cell in row):
            column = None
            continue
        if row[0].startswith("#"):
            continue
        if column is None:
            # the first line of a table is its header
            tables += 1
            column = row.index("_time") if "_time" in row else -1
            tables_with_time += column >= 0
            continue
        if column < 0 or column >= len(row) or not row[column]:
            continue
        try:
            timestamp = parse_timestamp(row[column])
        except ValueError:
            raise PluginError(UNKNOWN, "cannot parse the _time value %r" % row[column])
        if newest is None or timestamp > newest:
            newest = timestamp
    if tables and not tables_with_time:
        raise PluginError(UNKNOWN, "the query result has no _time column")
    return newest


def read_token(args):
    """The token from -f, $INFLUX_TOKEN_FILE, -t or $INFLUX_TOKEN, in this order."""
    path = args.token_file or os.environ.get("INFLUX_TOKEN_FILE")
    if path:
        try:
            with open(path) as handle:
                token = handle.read().strip()
        except OSError as err:
            raise PluginError(UNKNOWN, "cannot read token file %s: %s" % (path, err.strerror))
        if not token:
            raise PluginError(UNKNOWN, "token file %s is empty" % path)
        return token
    token = args.token or os.environ.get("INFLUX_TOKEN")
    if not token:
        raise PluginError(UNKNOWN, "no API token given, use -f/--token-file")
    return token


def api_error(body):
    """The message of an InfluxDB error answer, which is JSON like {"code":"invalid","message":"..."}."""
    try:
        data = json.loads(body.decode("utf-8", "replace"))
        return data.get("message") or data.get("code") or ""
    except (ValueError, AttributeError):
        return body.decode("utf-8", "replace").strip()[:200]


def query(args, token):
    scheme = "https" if args.ssl else "http"
    host = "[%s]" % args.host if ":" in args.host else args.host
    params = {"orgID": args.org_id} if args.org_id else {"org": args.org}
    url = "%s://%s:%d/api/v2/query?%s" % (scheme, host, args.port, urllib.parse.urlencode(params))
    request = urllib.request.Request(url, data=args.query.encode("utf-8"), method="POST", headers={
        "Authorization": "Token %s" % token,
        "Accept": "application/csv",
        "Content-Type": "application/vnd.flux",
        "User-Agent": "check_influxdb_data_age/%s" % VERSION,
    })
    if args.verbose:
        # never the token itself
        sys.stderr.write("POST %s\nAuthorization: Token ***\n%s\n" % (url, args.query))
    context = None
    if args.ssl:
        context = ssl.create_default_context(cafile=args.ca_cert)
        if args.insecure:
            context.check_hostname = False
            context.verify_mode = ssl.CERT_NONE
    try:
        with urllib.request.urlopen(request, timeout=args.timeout, context=context) as response:
            body = response.read()
    except urllib.error.HTTPError as err:
        message = api_error(err.read())
        detail = ": %s" % message if message else ""
        if err.code in (401, 403):
            raise PluginError(UNKNOWN, "InfluxDB rejected the token (HTTP %d)%s" % (err.code, detail))
        if err.code >= 500:
            raise PluginError(CRITICAL, "InfluxDB answered with HTTP %d%s" % (err.code, detail))
        raise PluginError(UNKNOWN, "InfluxDB answered with HTTP %d%s" % (err.code, detail))
    except (socket.timeout, TimeoutError):
        raise PluginError(CRITICAL, "no answer from %s:%d within %ss" % (args.host, args.port,
                                                                        format_number(args.timeout)))
    except (urllib.error.URLError, OSError) as err:
        reason = getattr(err, "reason", err)
        if isinstance(reason, (socket.timeout, TimeoutError)):
            raise PluginError(CRITICAL, "no answer from %s:%d within %ss" % (args.host, args.port,
                                                                            format_number(args.timeout)))
        raise PluginError(CRITICAL, "cannot query %s:%d: %s" % (args.host, args.port,
                                                                getattr(reason, "strerror", None) or reason))
    text = body.decode("utf-8", "replace")
    if args.verbose:
        sys.stderr.write(text + "\n")
    return text


def check(args, text, now):
    newest = newest_timestamp(text)
    if newest is None:
        return STATE_BY_NAME[args.empty_state], "the query returned no data", []

    age = round((now - newest).total_seconds(), 1)
    perf = [perfdata("age", age, "s", args.warning, args.critical)]
    when = newest.astimezone(datetime.timezone.utc).strftime("%Y-%m-%d %H:%M:%S UTC")
    if age < 0:
        # a timestamp in the future: the clock of the writer or of this host is off
        message = "newest entry is %s in the future (%s)" % (format_duration(-age), when)
        allowed = args.allow_timestamps_in_future
        if allowed is None:
            return WARNING, message + ", see --allow-timestamps-in-future", perf
        if -age > allowed:
            return WARNING, message + ", more than the allowed %s" % format_duration(allowed), perf
        # the data is current; the thresholds apply to old data only
        age = 0
    else:
        message = "newest entry is %s old (%s)" % (format_duration(age), when)

    if args.critical and args.critical.breached(age):
        state = CRITICAL
    elif args.warning and args.warning.breached(age):
        state = WARNING
    else:
        state = OK
    return state, message, perf


def duration_type(text):
    """A non-negative duration in seconds, with an optional unit s, m, h or d."""
    try:
        value = parse_number(text)
    except ValueError:
        raise argparse.ArgumentTypeError("invalid duration %r" % text)
    if value < 0:
        raise argparse.ArgumentTypeError("the duration must not be negative")
    return value


class ArgumentParser(argparse.ArgumentParser):
    """argparse exits 2 on a usage error, which Nagios reads as CRITICAL; exit UNKNOWN instead."""

    def error(self, message):
        self.print_usage(sys.stderr)
        self.exit(UNKNOWN, "%s: error: %s\n" % (self.prog, message))


def build_parser():
    parser = ArgumentParser(
        prog="check_influxdb_data_age",
        description="Nagios/Icinga plugin to check the age of the newest data an InfluxDB 2 Flux query returns, "
                    "via the query API (https://docs.influxdata.com/influxdb/v2/query-data/execute-queries/"
                    "influx-api/). The age is that of the newest _time over all tables and rows of the answer.",
        epilog="Thresholds take the monitoring plugins range format, in seconds; a unit s, m, h or d may "
               "follow each number, so the defaults -w 10m -c 30m alert when the newest entry is older than "
               "10 and 30 minutes.")
    parser.add_argument("-V", "--version", action="version", version="check_influxdb_data_age %s" % VERSION)
    parser.add_argument("-H", "--host", default="localhost", help="InfluxDB host (default: localhost)")
    parser.add_argument("-p", "--port", type=int, default=DEFAULT_PORT,
                        help="InfluxDB port (default: %d)" % DEFAULT_PORT)
    parser.add_argument("-S", "--ssl", "-ssl", action="store_true", help="use HTTPS instead of HTTP")
    parser.add_argument("--insecure", action="store_true", help="with --ssl, do not verify the certificate")
    parser.add_argument("--ca-cert", help="with --ssl, CA bundle to verify the certificate with")
    parser.add_argument("-T", "--timeout", type=float, default=DEFAULT_TIMEOUT,
                        help="timeout in seconds (default: %d)" % DEFAULT_TIMEOUT)
    token = parser.add_mutually_exclusive_group()
    token.add_argument("-f", "--token-file", help="file holding the API token (also: $INFLUX_TOKEN_FILE)")
    token.add_argument("-t", "--token", help="API token (also: $INFLUX_TOKEN); visible in the process list, "
                                             "prefer -f")
    org = parser.add_mutually_exclusive_group(required=True)
    org.add_argument("-o", "--org-id", help="ID of the organization")
    org.add_argument("--org", help="name of the organization")
    parser.add_argument("--query", required=True,
                        help="Flux query; the answer needs a _time column, e.g. "
                             "'from(bucket: \"sensors\") |> range(start: -1d) |> last()'")
    parser.add_argument("-w", "--warning", default=DEFAULT_WARNING,
                        help="warning range for the age of the newest entry (default: %s)" % DEFAULT_WARNING)
    parser.add_argument("-c", "--critical", default=DEFAULT_CRITICAL,
                        help="critical range for the age of the newest entry (default: %s)" % DEFAULT_CRITICAL)
    parser.add_argument("--allow-timestamps-in-future", nargs="?", type=duration_type, const=float("inf"),
                        metavar="DURATION",
                        help="do not report entries with a timestamp in the future, which are WARNING otherwise; "
                             "with DURATION, e.g. 30s, only up to that far in the future, to allow for clock "
                             "offsets")
    parser.add_argument("--empty-state", choices=("ok", "warning", "critical", "unknown"), default="critical",
                        help="state if the query returns no data, e.g. none within its range() (default: critical)")
    parser.add_argument("-v", "--verbose", action="store_true",
                        help="print the request, without the token, and the answer to stderr")
    return parser


def main():
    args = build_parser().parse_args()
    try:
        if args.timeout <= 0:
            raise PluginError(UNKNOWN, "-T must be above 0")
        args.warning = Range.parse(args.warning) if args.warning else None
        args.critical = Range.parse(args.critical) if args.critical else None
        token = read_token(args)
        text = query(args, token)
        state, message, perf = check(args, text, datetime.datetime.now(datetime.timezone.utc))
    except PluginError as err:
        print("INFLUXDB DATA AGE %s - %s" % (STATE_TEXT[err.state], err.message))
        return err.state
    except Exception as err:  # noqa: BLE001 - one line of UNKNOWN instead of a traceback
        print("INFLUXDB DATA AGE UNKNOWN - unexpected error: %s: %s" % (type(err).__name__, err))
        return UNKNOWN

    output = "INFLUXDB DATA AGE %s - %s" % (STATE_TEXT[state], message)
    if perf:
        output += " | " + " ".join(perf)
    print(output)
    return state


if __name__ == "__main__":
    try:
        sys.exit(main())
    except KeyboardInterrupt:
        sys.exit(UNKNOWN)
