#!/usr/bin/python3
# Nagios/Icinga plugin to check a STUN/TURN server such as coturn.
#
# Copyright (C) 2026 Thomas Wagner <wagner-thomas@gmx.at>
# SPDX-License-Identifier: GPL-2.0-or-later

import argparse
import base64
import datetime
import hashlib
import hmac
import os
import socket
import ssl
import struct
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 = 3478
DEFAULT_TLS_PORT = 5349
DEFAULT_TIMEOUT = 10
DEFAULT_SECRET_USER = "check_stun_turn"
SECRET_VALIDITY = 3600

MAGIC_COOKIE = 0x2112A442
BINDING = 0x001
ALLOCATE = 0x003
REFRESH = 0x004
SUCCESS = 0x100
ERROR = 0x110

MAPPED_ADDRESS = 0x0001
USERNAME = 0x0006
MESSAGE_INTEGRITY = 0x0008
ERROR_CODE = 0x0009
LIFETIME = 0x000D
REALM = 0x0014
NONCE = 0x0015
XOR_RELAYED_ADDRESS = 0x0016
REQUESTED_TRANSPORT = 0x0019
XOR_MAPPED_ADDRESS = 0x0020
SOFTWARE = 0x8022

TRANSPORT_UDP = 17
STALE_NONCE = 438


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_secret(value, path, variable):
    if value:
        return value
    path = path or os.environ.get(variable + "_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 %s: %s" % (path, err))
    return os.environ.get(variable)


# --------------------------------------------------------------------------
# STUN messages (RFC 8489) with the TURN methods of RFC 8656
# --------------------------------------------------------------------------
def attribute(kind, value):
    return struct.pack("!HH", kind, len(value)) + value + b"\0" * (-len(value) % 4)


def encode(method, attributes, transaction, key=None):
    body = b"".join(attribute(kind, value) for kind, value in attributes)
    if key is not None:
        # The length in the header covers MESSAGE-INTEGRITY itself when the HMAC is computed.
        header = struct.pack("!HHI", method, len(body) + 24, MAGIC_COOKIE) + transaction
        body += attribute(MESSAGE_INTEGRITY, hmac.new(key, header + body, hashlib.sha1).digest())
    return struct.pack("!HHI", method, len(body), MAGIC_COOKIE) + transaction + body


class Response:
    def __init__(self, data):
        if len(data) < 20:
            raise PluginError(CRITICAL, "answer too short for a STUN message")
        self.type, length, cookie = struct.unpack("!HHI", data[:8])
        if cookie != MAGIC_COOKIE or self.type & 0xC000:
            raise PluginError(CRITICAL, "answer is not a STUN message")
        self.transaction = data[8:20]
        self.attributes = {}
        position = 20
        while position + 4 <= min(len(data), 20 + length):
            kind, size = struct.unpack("!HH", data[position:position + 4])
            self.attributes.setdefault(kind, data[position + 4:position + 4 + size])
            position += 4 + size + (-size % 4)

    @property
    def success(self):
        return self.type & 0x0110 == SUCCESS

    @property
    def error(self):
        """(code, reason) of an error response."""
        value = self.attributes.get(ERROR_CODE, b"\0\0\0\0")
        return (value[2] & 0x07) * 100 + value[3], value[4:].decode("utf-8", "replace")

    def text(self, kind):
        value = self.attributes.get(kind)
        return value.decode("utf-8", "replace") if value is not None else None

    def address(self, kind):
        value = self.attributes.get(kind)
        if value is None or len(value) < 8:
            return None
        family, port = struct.unpack("!xBH", value[:4])
        raw = value[4:]
        if kind != MAPPED_ADDRESS:
            port ^= MAGIC_COOKIE >> 16
            mask = struct.pack("!I", MAGIC_COOKIE) + (self.transaction if family == 2 else b"")
            raw = bytes(a ^ b for a, b in zip(raw, mask))
        if family == 1:
            return "%s:%d" % (socket.inet_ntop(socket.AF_INET, raw[:4]), port)
        return "[%s]:%d" % (socket.inet_ntop(socket.AF_INET6, raw[:16]), port)


# --------------------------------------------------------------------------
# Transports: UDP with the retransmissions of RFC 8489, TCP and TLS
# --------------------------------------------------------------------------
class Transport:
    def __init__(self, args):
        self.args = args
        self.kind = args.transport
        self.certificate = None
        family = {4: socket.AF_INET, 6: socket.AF_INET6}.get(args.family, socket.AF_UNSPEC)
        kind = socket.SOCK_DGRAM if self.kind == "udp" else socket.SOCK_STREAM
        try:
            info = socket.getaddrinfo(args.hostname, args.port, family, kind)[0]
        except socket.gaierror as err:
            raise PluginError(CRITICAL, "cannot resolve %s: %s" % (args.hostname, err))
        self.peer = info[4][0]
        try:
            self.socket = socket.socket(info[0], info[1])
            self.socket.settimeout(args.timeout)
            self.socket.connect(info[4])
        except socket.timeout:
            raise PluginError(CRITICAL, "timeout after %gs connecting to %s port %d/%s"
                              % (args.timeout, args.hostname, args.port, self.kind))
        except OSError as err:
            raise PluginError(CRITICAL, "cannot connect to %s port %d/%s: %s"
                              % (args.hostname, args.port, self.kind, err.strerror or err))
        if self.kind == "tls":
            self._start_tls()
        self.buffer = b""

    def _start_tls(self):
        context = ssl.create_default_context()
        if self.args.insecure:
            context.check_hostname = False
            context.verify_mode = ssl.CERT_NONE
        try:
            self.socket = context.wrap_socket(self.socket, server_hostname=self.args.sni or self.args.hostname)
        except ssl.SSLCertVerificationError as err:
            raise PluginError(CRITICAL, "TLS certificate of %s is not valid: %s" % (self.args.hostname, err.verify_message))
        except (ssl.SSLError, OSError) as err:
            raise PluginError(CRITICAL, "TLS handshake with %s failed: %s" % (self.args.hostname, err))
        if not self.args.insecure:
            self.certificate = self.socket.getpeercert()

    def transact(self, data, transaction):
        """Send a request and return the response with its round trip time."""
        started = time.monotonic()
        deadline = started + self.args.timeout
        if self.kind == "udp":
            timeout = 0.5
            while True:
                self.socket.send(data)
                wait_until = min(deadline, time.monotonic() + timeout)
                while time.monotonic() < wait_until:
                    self.socket.settimeout(wait_until - time.monotonic())
                    try:
                        packet = self.socket.recv(65536)
                    except socket.timeout:
                        break
                    except OSError as err:
                        raise PluginError(CRITICAL, "no answer from %s port %d/udp: %s"
                                          % (self.args.hostname, self.args.port, err.strerror or err))
                    if packet[8:20] == transaction:
                        return Response(packet), time.monotonic() - started
                if time.monotonic() >= deadline:
                    raise PluginError(CRITICAL, "no answer from %s port %d/udp within %gs"
                                      % (self.args.hostname, self.args.port, self.args.timeout))
                timeout *= 2

        self.socket.sendall(data)
        while True:
            while len(self.buffer) < 20 or len(self.buffer) < 20 + struct.unpack("!H", self.buffer[2:4])[0]:
                self.socket.settimeout(max(deadline - time.monotonic(), 0.001))
                try:
                    chunk = self.socket.recv(65536)
                except socket.timeout:
                    raise PluginError(CRITICAL, "no answer from %s port %d/%s within %gs"
                                      % (self.args.hostname, self.args.port, self.kind, self.args.timeout))
                except OSError as err:
                    raise PluginError(CRITICAL, "connection to %s port %d/%s failed: %s"
                                      % (self.args.hostname, self.args.port, self.kind, err))
                if not chunk:
                    raise PluginError(CRITICAL, "%s closed the connection" % self.args.hostname)
                self.buffer += chunk
            size = 20 + struct.unpack("!H", self.buffer[2:4])[0]
            packet, self.buffer = self.buffer[:size], self.buffer[size:]
            if packet[8:20] == transaction:
                return Response(packet), time.monotonic() - started

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


def certificate_state(args, transport, state, parts, perf):
    """Add the expiry of a validated TLS certificate to the result."""
    cert = transport.certificate
    if not cert or "notAfter" not in cert:
        return state
    expires = datetime.datetime.fromtimestamp(ssl.cert_time_to_seconds(cert["notAfter"]), datetime.timezone.utc)
    days = (expires - datetime.datetime.now(datetime.timezone.utc)).total_seconds() / 86400
    if days < args.cert_critical:
        state = max(state, CRITICAL)
    elif days < args.cert_warning:
        state = max(state, WARNING)
    parts.append("certificate valid for %.0f days (until %s)" % (days, expires.strftime("%Y-%m-%d")))
    perf.append("cert_days=%s;%d;%d" % (format_number(round(days, 1)), args.cert_warning, args.cert_critical))
    return state


def server_name(response):
    return response.text(SOFTWARE) or "server"


def check_stun(args):
    transport = Transport(args)
    try:
        transaction = os.urandom(12)
        response, elapsed = transport.transact(encode(BINDING, [], transaction), transaction)
    finally:
        transport.close()

    if not response.success:
        code, reason = response.error
        raise PluginError(CRITICAL, "%s refused the binding request: %d %s" % (server_name(response), code, reason))
    mapped = response.address(XOR_MAPPED_ADDRESS) or response.address(MAPPED_ADDRESS)
    if mapped is None:
        raise PluginError(CRITICAL, "%s answered without a mapped address" % server_name(response))

    state = evaluate(args, elapsed)
    parts = ["%s answered over %s in %.3fs, mapped address %s" % (server_name(response), args.transport.upper(),
                                                                  elapsed, mapped)]
    perf = [perfdata("time", round(elapsed, 3), "s", args.warning, args.critical, 0)]
    state = certificate_state(args, transport, state, parts, perf)
    return state, ", ".join(parts), perf


def credentials(args):
    """User name and password: long-term credentials, or derived from the shared secret of the TURN REST API."""
    secret = read_secret(args.secret, args.secret_file, "TURN_SECRET")
    if secret:
        username = "%d:%s" % (int(time.time()) + SECRET_VALIDITY, args.username or DEFAULT_SECRET_USER)
        password = base64.b64encode(hmac.new(secret.encode(), username.encode(), hashlib.sha1).digest()).decode()
        return username, password
    password = read_secret(args.password, args.password_file, "TURN_PASSWORD")
    if args.username and password is None:
        raise PluginError(UNKNOWN, "a user name needs a password, use -P, -f or $TURN_PASSWORD")
    if password is not None and not args.username:
        raise PluginError(UNKNOWN, "a password needs a user name, use -u")
    return (args.username, password) if args.username else (None, None)


def authenticated(transport, method, attributes, username, key, challenge):
    """Send an authenticated request, answering one stale nonce with the new one."""
    realm, nonce = challenge
    for _ in range(2):
        transaction = os.urandom(12)
        signed = attributes + [(USERNAME, username.encode()), (REALM, realm.encode()), (NONCE, nonce.encode())]
        response, elapsed = transport.transact(encode(method, signed, transaction, key), transaction)
        if response.success or response.error[0] != STALE_NONCE:
            return response, elapsed
        nonce = response.text(NONCE) or nonce
    return response, elapsed


def check_turn(args):
    username, password = credentials(args)
    transport = Transport(args)
    try:
        transaction = os.urandom(12)
        requested = [(REQUESTED_TRANSPORT, struct.pack("!B3x", TRANSPORT_UDP))]
        first, elapsed = transport.transact(encode(ALLOCATE, requested, transaction), transaction)
        name = server_name(first)
        perf = []
        parts = []
        state = OK

        if first.success:
            state = WARNING
            parts.append("%s allocates relays without authentication, it is an open relay" % name)
            release = os.urandom(12)
            transport.transact(encode(REFRESH, [(LIFETIME, struct.pack("!I", 0))], release), release)
        else:
            code, reason = first.error
            realm, nonce = first.text(REALM), first.text(NONCE)
            if code != 401 or not realm or not nonce:
                raise PluginError(CRITICAL, "%s refused the allocation: %d %s" % (name, code, reason))
            if username is None:
                parts.append("%s asks for authentication (realm %s), no credentials given so no relay was allocated"
                             % (name, realm))
            else:
                key = hashlib.md5(("%s:%s:%s" % (username, realm, password)).encode("utf-8")).digest()
                response, second = authenticated(transport, ALLOCATE, requested, username, key, (realm, nonce))
                elapsed += second
                if not response.success:
                    code, reason = response.error
                    if code == 401:
                        raise PluginError(UNKNOWN, "%s rejected the credentials for realm %s" % (name, realm))
                    raise PluginError(CRITICAL, "%s refused the allocation: %d %s" % (name, code, reason))
                relayed = response.address(XOR_RELAYED_ADDRESS) or "?"
                lifetime = struct.unpack("!I", response.attributes.get(LIFETIME, b"\0\0\0\0"))[0]
                parts.append("%s allocated relay %s for %ds in realm %s" % (name, relayed, lifetime, realm))
                nonce = response.text(NONCE) or nonce
                released, _ = authenticated(transport, REFRESH, [(LIFETIME, struct.pack("!I", 0))], username, key,
                                            (realm, nonce))
                if not released.success:
                    state = WARNING
                    parts.append("releasing it failed: %d %s" % released.error)
    finally:
        transport.close()

    state = max(state, evaluate(args, elapsed))
    parts[0] = parts[0] + ", %.3fs over %s" % (elapsed, args.transport.upper())
    perf.append(perfdata("time", round(elapsed, 3), "s", args.warning, args.critical, 0))
    state = certificate_state(args, transport, state, parts, perf)
    return state, ", ".join(parts), perf


MODES = {
    "stun": (check_stun, "1", "3"),
    "turn": (check_turn, "2", "5"),
}


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_stun_turn",
                            description="Nagios/Icinga plugin to check a STUN/TURN server such as coturn.")
    parser.add_argument("-V", "--version", action="version", version="check_stun_turn %s" % VERSION)

    common = ArgumentParser(add_help=False)
    common.add_argument("-H", "--hostname", required=True, help="host name or address of the server")
    common.add_argument("-p", "--port", type=int,
                        help="port (default: %d, %d with --transport tls)" % (DEFAULT_PORT, DEFAULT_TLS_PORT))
    common.add_argument("-T", "--transport", choices=("udp", "tcp", "tls"), default="udp",
                        help="transport to the server (default: udp)")
    family = common.add_mutually_exclusive_group()
    family.add_argument("-4", dest="family", action="store_const", const=4, help="use IPv4")
    family.add_argument("-6", dest="family", action="store_const", const=6, help="use IPv6")
    common.add_argument("-t", "--timeout", type=float, default=DEFAULT_TIMEOUT,
                        help="timeout in seconds (default: %d)" % DEFAULT_TIMEOUT)
    common.add_argument("-w", "--warning", help="warning threshold for the response time in seconds, as a range")
    common.add_argument("-c", "--critical", help="critical threshold for the response time in seconds, as a range")
    common.add_argument("-k", "--insecure", action="store_true", help="do not verify the TLS certificate")
    common.add_argument("--sni", help="server name for TLS (default: the host name)")
    common.add_argument("--cert-warning", type=int, default=30, metavar="DAYS",
                        help="warn when the TLS certificate expires in fewer days (default: 30)")
    common.add_argument("--cert-critical", type=int, default=14, metavar="DAYS",
                        help="alert when the TLS certificate expires in fewer days (default: 14)")

    subparsers = parser.add_subparsers(dest="mode", required=True, metavar="MODE")
    subparsers.add_parser("stun", parents=[common],
                          help="send a STUN binding request (defaults to -w 1 -c 3)",
                          description="Send a STUN binding request and check the answer and the mapped address. "
                                      "The thresholds apply to the response time (defaults to -w 1 -c 3).")
    turn = subparsers.add_parser(
        "turn", parents=[common], help="request a TURN allocation (defaults to -w 2 -c 5)",
        description="Request a TURN allocation. Without credentials the server has to answer with an "
                    "authentication challenge; with credentials a relay is allocated and released again. The "
                    "thresholds apply to the total response time (defaults to -w 2 -c 5).")
    turn.add_argument("-u", "--username", help="TURN user name; with a shared secret, the name the credentials "
                                               "are issued for (default: %s)" % DEFAULT_SECRET_USER)
    turn.add_argument("-P", "--password", help="TURN password; visible in the process list, prefer -f")
    turn.add_argument("-f", "--password-file", help="file holding the TURN password")
    turn.add_argument("-s", "--secret", help="shared secret of the TURN REST API (coturn: static-auth-secret); visible in the process list, prefer -S")
    turn.add_argument("-S", "--secret-file", help="file holding the shared secret")
    return parser


def main():
    args = build_parser().parse_args()
    if args.port is None:
        args.port = DEFAULT_TLS_PORT if args.transport == "tls" else DEFAULT_PORT
    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, struct.error) as err:
        print("%s UNKNOWN - unexpected answer from the server: %r" % (args.mode.upper(), err))
        return UNKNOWN
    except PluginError as err:
        print("%s %s - %s" % (args.mode.upper(), STATE_TEXT[err.state], err.message))
        return err.state

    output = "%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)
