#!/usr/bin/python3
# Nagios/Icinga plugin to check that an NTRIP caster serves current GNSS correction data.
#
# Copyright (C) 2026 Thomas Wagner <thomas.wagner@jku.at>
# SPDX-License-Identifier: GPL-2.0-or-later

import argparse
import base64
import datetime
import email.utils
import os
import socket
import sys
import time

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

STATE_TEXT = {OK: "OK", WARNING: "WARNING", CRITICAL: "CRITICAL", UNKNOWN: "UNKNOWN"}

VERSION = "0.1"
DEFAULT_PORT = 2101
DEFAULT_TIMEOUT = 10
DEFAULT_DURATION = 5
DEFAULT_LEAP_SECONDS = 18
USER_AGENT = "NTRIP check_ntrip_caster/%s" % VERSION

GPS_EPOCH_UNIX = 315964800
WEEK_MS = 7 * 86400 * 1000
DAY_MS = 86400 * 1000
BEIDOU_OFFSET_MS = 14000
GLONASS_OFFSET_MS = 3 * 3600 * 1000


def _types(*spans):
    return {t for first, last in spans for t in range(first, last + 1)}


# RTCM3 messages that carry an observation epoch at bit 24 of the payload.
EPOCH_TOW = {"GPS": _types((1001, 1004), (1071, 1077)), "Galileo": _types((1091, 1097)),
             "SBAS": _types((1101, 1107)), "QZSS": _types((1111, 1117)), "NavIC": _types((1131, 1137))}
EPOCH_BEIDOU = _types((1121, 1127))
EPOCH_GLONASS_MSM = _types((1081, 1087))
EPOCH_GLONASS_LEGACY = _types((1009, 1012))


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


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

    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 float(low)
            end = float("inf") if high == "" else float(high)
        except ValueError:
            raise PluginError(UNKNOWN, "invalid threshold %r" % raw)

        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 isinstance(value, int) or float(value).is_integer():
        return "%d" % value
    return ("%.3f" % value).rstrip("0").rstrip(".")


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


def evaluate(args, value):
    if args.critical and args.critical.breached(value):
        return CRITICAL
    if args.warning and args.warning.breached(value):
        return WARNING
    return OK


def read_password(args):
    if args.password:
        return args.password

    path = args.password_file or os.environ.get("NTRIP_PASSWORD_FILE")
    if path:
        try:
            with open(os.path.expanduser(path), "r") as handle:
                return handle.read().strip()
        except OSError as err:
            raise PluginError(UNKNOWN, "cannot read password file: %s" % err)

    return os.environ.get("NTRIP_PASSWORD", "")


# --------------------------------------------------------------------------
# NTRIP 1.0 over a plain socket: casters answer "SOURCETABLE 200 OK" or
# "ICY 200 OK", which no HTTP library accepts as a status line.
# --------------------------------------------------------------------------
class Connection:
    def __init__(self, args, path, authenticate):
        self.args = args
        self.buffer = b""
        self.closed = False
        lines = ["GET /%s HTTP/1.0" % path, "User-Agent: %s" % USER_AGENT,
                 "Host: %s:%d" % (args.hostname, args.port)]
        if authenticate and args.username:
            credentials = "%s:%s" % (args.username, read_password(args))
            lines.append("Authorization: Basic %s" % base64.b64encode(credentials.encode("utf-8")).decode("ascii"))
        self.started = time.monotonic()
        try:
            self.socket = socket.create_connection((args.hostname, args.port), timeout=args.timeout)
            self.socket.sendall(("\r\n".join(lines) + "\r\n\r\n").encode("ascii"))
        except socket.timeout:
            raise PluginError(CRITICAL, "timeout after %gs connecting to %s:%d" % (args.timeout, args.hostname, args.port))
        except OSError as err:
            raise PluginError(CRITICAL, "cannot connect to %s:%d: %s" % (args.hostname, args.port, err))

    def fill(self, deadline=None):
        """Read more data; False once the caster closed the connection or the deadline passed."""
        if self.closed:
            return False
        timeout = self.args.timeout
        if deadline is not None:
            timeout = min(timeout, deadline - time.monotonic())
            if timeout <= 0:
                return False
        self.socket.settimeout(timeout)
        try:
            chunk = self.socket.recv(65536)
        except socket.timeout:
            if deadline is not None and time.monotonic() >= deadline:
                return False
            raise PluginError(CRITICAL, "no data from %s:%d for %gs" % (self.args.hostname, self.args.port, timeout))
        except OSError as err:
            raise PluginError(CRITICAL, "connection to %s:%d failed: %s" % (self.args.hostname, self.args.port, err))
        if not chunk:
            self.closed = True
            return False
        self.buffer += chunk
        return True

    def read_line(self):
        while b"\r\n" not in self.buffer:
            if len(self.buffer) > 8192 or not self.fill():
                raise PluginError(CRITICAL, "no valid answer from %s:%d, got %r"
                                  % (self.args.hostname, self.args.port, self.buffer[:60]))
        line, self.buffer = self.buffer.split(b"\r\n", 1)
        return line.decode("latin-1").strip()

    def read_headers(self):
        headers = {}
        while True:
            line = self.read_line()
            if not line:
                return headers
            name, _, value = line.partition(":")
            headers[name.strip().lower()] = value.strip()

    def close(self):
        self.socket.close()


def parse_date(text):
    """Caster time from the Date header: RFC 1123, or the YYYY/MM/DD HH:MM:SS UTC that RTKLIB sends."""
    try:
        stamp = email.utils.parsedate_to_datetime(text)
    except (TypeError, ValueError):
        stamp = None
    if stamp is None:
        try:
            stamp = datetime.datetime.strptime(text.replace(" UTC", "").replace(" GMT", ""), "%Y/%m/%d %H:%M:%S")
        except ValueError:
            return None
    if stamp.tzinfo is None:
        stamp = stamp.replace(tzinfo=datetime.timezone.utc)
    return stamp


def read_sourcetable(connection):
    lines = []
    while True:
        while b"\r\n" not in connection.buffer and b"\n" not in connection.buffer:
            if not connection.fill():
                if connection.buffer:
                    lines.append(connection.buffer.decode("latin-1").strip())
                    connection.buffer = b""
                return lines, False
        separator = b"\r\n" if b"\r\n" in connection.buffer else b"\n"
        line, connection.buffer = connection.buffer.split(separator, 1)
        line = line.decode("latin-1").strip()
        if line == "ENDSOURCETABLE":
            return lines, True
        if line:
            lines.append(line)


def check_sourcetable(args):
    connection = Connection(args, "", authenticate=False)
    try:
        status = connection.read_line()
        if not status.startswith("SOURCETABLE 200"):
            raise PluginError(CRITICAL, "caster answered %r instead of a sourcetable" % status)
        received = time.time()
        headers = connection.read_headers()
        lines, complete = read_sourcetable(connection)
        elapsed = time.monotonic() - connection.started
    finally:
        connection.close()

    if not complete:
        raise PluginError(CRITICAL, "sourcetable ends without ENDSOURCETABLE, %d lines received" % len(lines))

    streams = [line.split(";")[1] for line in lines if line.startswith("STR;") and line.count(";") >= 1]
    state = OK
    problems = []

    if args.mountpoint and args.mountpoint not in streams:
        state = CRITICAL
        problems.append("mountpoint %s is not listed" % args.mountpoint)

    perf = [perfdata("time", round(elapsed, 3), "s", None, None, 0), perfdata("streams", len(streams), "", None, None, 0)]
    date = parse_date(headers["date"]) if "date" in headers else None
    if date is None:
        state = max(state, WARNING)
        problems.append("no readable Date header" if "date" not in headers else "unreadable Date %r" % headers["date"])
        clock = ""
    else:
        # The Date header is truncated to whole seconds; +0.5 s centres that error.
        offset = date.timestamp() + 0.5 - received
        state = max(state, evaluate(args, abs(offset)))
        clock = ", caster clock %+.1fs off" % offset
        perf.append(perfdata("clock_offset", round(offset, 1), "s", args.warning, args.critical))

    listed = ", ".join(streams) if streams else "no streams"
    message = "%s lists %d stream(s) (%s)%s" % (headers.get("server", "caster"), len(streams), listed, clock)
    if problems:
        message = "; ".join(problems) + ", " + message
    return state, message, perf


# --------------------------------------------------------------------------
# RTCM3 framing: 0xD3, 6 reserved bits, 10 bit length, payload, CRC-24Q
# --------------------------------------------------------------------------
def crc24q(data):
    crc = 0
    for byte in data:
        crc ^= byte << 16
        for _ in range(8):
            crc <<= 1
            if crc & 0x1000000:
                crc ^= 0x1864CFB
    return crc & 0xFFFFFF


def bits(payload, position, count):
    value = 0
    for index in range(position, position + count):
        value = (value << 1) | ((payload[index // 8] >> (7 - index % 8)) & 1)
    return value


def signed_age(now_ms, epoch_ms, period_ms):
    return ((now_ms - epoch_ms + period_ms // 2) % period_ms - period_ms // 2) / 1000.0


def epoch_age(message_type, payload, now, leap_seconds):
    """Age in seconds of the observation epoch of a message, with its constellation, or None."""
    gps_ms = int((now - GPS_EPOCH_UNIX + leap_seconds) * 1000) % WEEK_MS
    for system, types in EPOCH_TOW.items():
        if message_type in types:
            return signed_age(gps_ms, bits(payload, 24, 30), WEEK_MS), system
    if message_type in EPOCH_BEIDOU:
        return signed_age(gps_ms, bits(payload, 24, 30) + BEIDOU_OFFSET_MS, WEEK_MS), "BeiDou"
    if message_type in EPOCH_GLONASS_MSM or message_type in EPOCH_GLONASS_LEGACY:
        tod = bits(payload, 27, 27) if message_type in EPOCH_GLONASS_MSM else bits(payload, 24, 27)
        utc_ms = int(now * 1000) % DAY_MS
        return signed_age(utc_ms, (tod - GLONASS_OFFSET_MS) % DAY_MS, DAY_MS), "GLONASS"
    return None


def check_stream(args):
    connection = Connection(args, args.mountpoint, authenticate=True)
    types = {}
    frames = crc_errors = 0
    newest = None
    first_frame = None
    synced = False
    try:
        status = connection.read_line()
        if status.startswith("SOURCETABLE"):
            raise PluginError(CRITICAL, "caster does not serve mountpoint %s, it answered with its sourcetable"
                              % args.mountpoint)
        if " 401" in status:
            raise PluginError(UNKNOWN, "credentials rejected for mountpoint %s (%s)" % (args.mountpoint, status))
        if not (status.startswith("ICY 200") or (status.startswith("HTTP/") and " 200" in status)):
            raise PluginError(CRITICAL, "caster answered %r for mountpoint %s" % (status, args.mountpoint))
        if status.startswith("HTTP/"):
            connection.read_headers()

        received = 0
        deadline = time.monotonic() + args.duration
        while True:
            buffer = connection.buffer
            start = buffer.find(b"\xd3")
            if start < 0 or len(buffer) < start + 6:
                if not connection.fill(deadline):
                    break
                continue
            length = ((buffer[start + 1] & 0x03) << 8) | buffer[start + 2]
            if len(buffer) < start + 6 + length:
                if not connection.fill(deadline):
                    break
                continue
            frame = buffer[start:start + 6 + length]
            if crc24q(frame[:-3]) != int.from_bytes(frame[-3:], "big"):
                # Before the first valid frame this is just a 0xD3 inside other data.
                if synced:
                    crc_errors += 1
                    synced = False
                connection.buffer = buffer[start + 1:]
                continue
            synced = True
            frames += 1
            received += len(frame)
            if first_frame is None:
                first_frame = time.monotonic() - connection.started
            payload = frame[3:-3]
            if length >= 2:
                message_type = bits(payload, 0, 12)
                types[message_type] = types.get(message_type, 0) + 1
                if length >= 8:
                    age = epoch_age(message_type, payload, time.time(), args.leap_seconds)
                    if age is not None and (newest is None or abs(age[0]) < abs(newest[0])):
                        newest = age
            connection.buffer = buffer[start + 6 + length:]
        listened = time.monotonic() - connection.started
        closed = connection.closed
    finally:
        connection.close()

    if not frames:
        raise PluginError(CRITICAL, "no RTCM3 data from mountpoint %s within %.1fs" % (args.mountpoint, listened))

    perf = [
        perfdata("frames", frames, "", None, None, 0),
        perfdata("bytes", received, "B", None, None, 0),
        perfdata("crc_errors", crc_errors, "", None, None, 0),
        perfdata("first_frame", round(first_frame, 3), "s", None, None, 0),
    ]
    summary = "%d RTCM3 frames in %.1fs, messages %s" % (frames, listened, " ".join(str(t) for t in sorted(types)))

    if newest is None:
        raise PluginError(UNKNOWN, "mountpoint %s sends no observation message with an epoch time, so its age "
                                   "cannot be told; %s" % (args.mountpoint, summary))

    age, system = newest
    state = evaluate(args, abs(age))
    problems = []
    if crc_errors:
        state = max(state, WARNING)
        problems.append("%d frame(s) with a wrong checksum" % crc_errors)
    if closed:
        state = max(state, WARNING)
        problems.append("caster closed the stream after %.1fs" % listened)

    message = "%s: newest epoch %.2fs old (%s), %s" % (args.mountpoint, age, system, summary)
    if problems:
        message = "; ".join(problems) + ", " + message
    perf.insert(0, perfdata("age", round(age, 3), "s", args.warning, args.critical))
    return state, message, perf


MODES = {
    "sourcetable": (check_sourcetable, "check the sourcetable, the caster clock and optionally a mountpoint",
                    "10", "60"),
    "stream": (check_stream, "check that a mountpoint streams valid RTCM3 data with a current epoch", "3", "10"),
}


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)
        sys.stderr.write("%s: error: %s\n" % (self.prog, message))
        sys.exit(UNKNOWN)


def build_parser():
    parser = ArgumentParser(
        prog="check_ntrip_caster",
        description="Nagios/Icinga plugin to check that an NTRIP caster serves current GNSS correction data.",
    )
    parser.add_argument("-V", "--version", action="version", version="check_ntrip_caster %s" % VERSION)

    common = ArgumentParser(add_help=False)
    common.add_argument("-H", "--hostname", required=True, help="host name or address of the caster")
    common.add_argument("-p", "--port", type=int, default=DEFAULT_PORT,
                        help="port of the caster (default: %d)" % DEFAULT_PORT)
    common.add_argument("-t", "--timeout", type=float, default=DEFAULT_TIMEOUT,
                        help="timeout in seconds for connecting and between data (default: %d)" % DEFAULT_TIMEOUT)
    common.add_argument("-w", "--warning", help="warning threshold as a range")
    common.add_argument("-c", "--critical", help="critical threshold as a range")

    subparsers = parser.add_subparsers(dest="mode", required=True, metavar="MODE")

    source = subparsers.add_parser(
        "sourcetable", parents=[common],
        help="check the sourcetable, the caster clock and optionally a mountpoint (defaults to -w 10 -c 60)",
        description="Fetch the sourcetable. The thresholds apply to the difference between the caster's Date "
                    "header and the local clock in seconds (defaults to -w 10 -c 60).")
    source.add_argument("-m", "--mountpoint", help="mountpoint that has to be listed in the sourcetable")

    stream = subparsers.add_parser(
        "stream", parents=[common],
        help="check that a mountpoint streams valid RTCM3 data with a current epoch (defaults to -w 3 -c 10)",
        description="Log in to a mountpoint and read RTCM3 data. The thresholds apply to the age in seconds of "
                    "the newest observation epoch, compared with the local clock (defaults to -w 3 -c 10).")
    stream.add_argument("-m", "--mountpoint", required=True, help="mountpoint to read")
    stream.add_argument("-u", "--username", help="user name for the mountpoint")
    stream.add_argument("-P", "--password", help="password; visible in the process list, prefer -f")
    stream.add_argument("-f", "--password-file", help="file holding the password")
    stream.add_argument("-d", "--duration", type=float, default=DEFAULT_DURATION,
                        help="seconds to read the stream (default: %d)" % DEFAULT_DURATION)
    stream.add_argument("--leap-seconds", type=int, default=DEFAULT_LEAP_SECONDS,
                        help="GPS minus UTC in seconds (default: %d)" % DEFAULT_LEAP_SECONDS)

    return parser


def main():
    args = build_parser().parse_args()
    handler, _, default_warning, default_critical = MODES[args.mode]
    warning = args.warning if args.warning is not None else default_warning
    critical = args.critical if args.critical is not None else default_critical

    try:
        args.warning = Range.parse(warning) if warning else None
        args.critical = Range.parse(critical) if critical else None
        state, message, perf = handler(args)
    except (KeyError, TypeError, ValueError, IndexError) as err:
        print("NTRIP %s UNKNOWN - unexpected answer from the caster: %r" % (args.mode.upper(), err))
        return UNKNOWN
    except PluginError as err:
        print("NTRIP %s %s - %s" % (args.mode.upper(), STATE_TEXT[err.state], err.message))
        return err.state

    output = "NTRIP %s %s - %s" % (args.mode.upper(), 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)
