#!/usr/bin/python3
# Nagios/Icinga plugin to check a GitLab instance via its REST API.
#
# Copyright (c) 2026 Thomas Wagner <wagner-thomas@gmx.at>
# SPDX-License-Identifier: GPL-2.0-or-later

import argparse
import datetime
import fcntl
import json
import os
import socket
import ssl
import stat
import sys
import tempfile
import time
import urllib.error
import urllib.request

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

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

VERSION = "0.2"
DEFAULT_PORT = 443
DEFAULT_TIMEOUT = 10
DEFAULT_LIFETIME = 90
DEFAULT_ROTATE_BEFORE = 30

MIGRATION_DONE = ("finished", "finalized")


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):
        """Range as perfdata carries it: same unit as the value."""
        if self.start == 0 and self.end != float("inf"):
            body = format_number(self.end)
        else:
            low = "~" if self.start == float("-inf") else format_number(self.start)
            high = "" if self.end == float("inf") else format_number(self.end)
            body = "%s:%s" % (low, high)
        return ("@" if self.inside else "") + body

    @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_token(args):
    if args.token:
        return args.token.strip()

    path = args.token_file or os.environ.get("GITLAB_TOKEN_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 token file: %s" % err)

    token = os.environ.get("GITLAB_TOKEN")
    if token:
        return token.strip()

    raise PluginError(UNKNOWN, "no access token given, use -T, -f or $GITLAB_TOKEN")


def http_request(args, path, token, method="GET", payload=None):
    scheme = "https" if args.ssl else "http"
    url = "%s://%s:%d/api/v4%s" % (scheme, args.hostname, args.port, path)

    headers = {"Accept": "application/json", "PRIVATE-TOKEN": token}
    data = None
    if payload is not None:
        data = json.dumps(payload).encode("utf-8")
        headers["Content-Type"] = "application/json"
    request = urllib.request.Request(url, data=data, headers=headers, method=method)

    context = None
    if args.ssl:
        context = ssl.create_default_context()
        if args.insecure:
            context.check_hostname = False
            context.verify_mode = ssl.CERT_NONE

    started = time.monotonic()
    with urllib.request.urlopen(request, timeout=args.timeout, context=context) as response:
        body = response.read()
    return json.loads(body.decode("utf-8")), time.monotonic() - started


def api_get(args, path, token=None):
    url = "%s://%s:%d/api/v4%s" % ("https" if args.ssl else "http", args.hostname, args.port, path)
    try:
        return http_request(args, path, token or read_token(args))
    except urllib.error.HTTPError as err:
        if err.code == 401:
            raise PluginError(UNKNOWN, "access token rejected (HTTP 401), is it expired or revoked?")
        if err.code == 403:
            raise PluginError(UNKNOWN, "access token may not read %s (HTTP 403), it needs the read_api "
                                       "scope and, for this mode, an administrator" % path)
        if err.code >= 500:
            raise PluginError(CRITICAL, "server error on %s (HTTP %d)" % (path, err.code))
        raise PluginError(UNKNOWN, "unexpected HTTP %d on %s" % (err.code, path))
    except urllib.error.URLError as err:
        raise PluginError(CRITICAL, "cannot reach %s: %s" % (url, err.reason))
    except socket.timeout:
        raise PluginError(CRITICAL, "timeout after %gs while reading %s" % (args.timeout, url))
    except (UnicodeDecodeError, ValueError):
        raise PluginError(UNKNOWN, "no valid JSON returned by %s, is this a GitLab server?" % path)


class TokenFile:
    """Token file that can be replaced safely while other checks read it."""

    def __init__(self, path):
        self.path = os.path.realpath(os.path.expanduser(path))
        self.directory = os.path.dirname(self.path)
        self.lock = None
        self.pending = None

    def __enter__(self):
        try:
            self.lock = open(self.path + ".lock", "a")
            fcntl.flock(self.lock, fcntl.LOCK_EX | fcntl.LOCK_NB)
        except BlockingIOError:
            raise PluginError(UNKNOWN, "another rotation of %s is running" % self.path)
        except OSError as err:
            raise PluginError(UNKNOWN, "cannot lock %s: %s" % (self.path, err))
        return self

    def __exit__(self, *exc):
        self.discard()
        self.lock.close()

    def read(self):
        try:
            with open(self.path, "r") as handle:
                return handle.read().strip()
        except OSError as err:
            raise PluginError(UNKNOWN, "cannot read token file: %s" % err)

    def prepare(self):
        """Runs before the rotation request, so a file that cannot be stored stops it."""
        try:
            info = os.stat(self.path)
            fd, name = tempfile.mkstemp(prefix=".%s." % os.path.basename(self.path), dir=self.directory)
        except OSError as err:
            raise PluginError(UNKNOWN, "cannot create a replacement for %s: %s" % (self.path, err))
        self.pending = (fd, name)
        try:
            os.fchown(fd, info.st_uid, info.st_gid)
            os.fchmod(fd, stat.S_IMODE(info.st_mode))
        except OSError as err:
            raise PluginError(UNKNOWN, "cannot keep owner and mode of %s (%s), rotate as its owner or as root"
                                       % (self.path, err))

    def commit(self, secret):
        fd, name = self.pending
        self.pending = None
        with os.fdopen(fd, "w") as handle:
            handle.write(secret)
            handle.flush()
            os.fsync(handle.fileno())
        os.replace(name, self.path)
        directory = os.open(self.directory, os.O_RDONLY)
        try:
            os.fsync(directory)
        finally:
            os.close(directory)

    def discard(self):
        if self.pending:
            fd, name = self.pending
            self.pending = None
            os.close(fd)
            os.unlink(name)


def check_health(args):
    data, elapsed = api_get(args, "/version")

    state = evaluate(args, elapsed)
    message = "GitLab %s answered in %.3fs" % (data["version"], elapsed)
    perf = [perfdata("time", round(elapsed, 3), "s", args.warning, args.critical, 0)]
    return state, message, perf


def check_token(args):
    data, _ = api_get(args, "/personal_access_tokens/self")
    name = data.get("name", "?")

    if data.get("revoked") or not data.get("active", True):
        raise PluginError(CRITICAL, "access token %r is revoked or inactive" % name)

    expires = data.get("expires_at")
    if not expires:
        return OK, "access token %r never expires" % name, []

    today = datetime.datetime.now(datetime.timezone.utc).date()
    days = (datetime.date.fromisoformat(expires[:10]) - today).days

    state = evaluate(args, days)
    message = "access token %r expires on %s, in %d day%s" % (name, expires[:10], days, "" if days == 1 else "s")
    perf = [perfdata("days_left", days, "", args.warning, args.critical, 0)]
    return state, message, perf


def check_stats(args):
    data, _ = api_get(args, "/application/statistics")
    stats = {key: int(value) for key, value in data.items()}

    active = stats["active_users"]
    state = evaluate(args, active)
    message = "%d active of %d users, %d projects in %d groups, %d issues, %d merge requests" % (
        active, stats["users"], stats["projects"], stats["groups"], stats["issues"], stats["merge_requests"],
    )
    perf = [perfdata("active_users", active, "", args.warning, args.critical, 0, stats["users"])]
    for key in ("users", "projects", "groups", "issues", "merge_requests", "notes",
                "forks", "snippets", "milestones", "ssh_keys"):
        if key in stats:
            perf.append(perfdata(key, stats[key], "", None, None, 0))
    return state, message, perf


def check_sidekiq(args):
    data, _ = api_get(args, "/sidekiq/compound_metrics")
    jobs = data["jobs"]
    processes = data["processes"]
    queues = data["queues"]

    if not processes:
        raise PluginError(CRITICAL, "no Sidekiq process is running, background jobs are not processed")

    backlog = int(jobs["enqueued"])
    latency = max((float(queue["latency"]) for queue in queues.values()), default=0.0)
    concurrency = sum(int(process["concurrency"]) for process in processes)
    busy = sum(int(process["busy"]) for process in processes)

    state = evaluate(args, backlog)
    message = "%d jobs enqueued, %d of %d threads busy in %d processes, highest queue latency %.1fs" % (
        backlog, busy, concurrency, len(processes), latency,
    )
    perf = [
        perfdata("enqueued", backlog, "", args.warning, args.critical, 0),
        perfdata("latency", round(latency, 3), "s", None, None, 0),
        perfdata("busy", busy, "", None, None, 0, concurrency),
        perfdata("processes", len(processes), "", None, None, 0),
        perfdata("dead", int(jobs["dead"]), "", None, None, 0),
        perfdata("failed", int(jobs["failed"]), "c", None, None, 0),
        perfdata("processed", int(jobs["processed"]), "c", None, None, 0),
    ]
    return state, message, perf


def check_migrations(args):
    data, _ = api_get(args, "/admin/batched_background_migrations")

    failed = [m["job_class_name"] for m in data if m["status"] == "failed"]
    pending = [m for m in data if m["status"] not in MIGRATION_DONE and m["status"] != "failed"]

    perf = [
        perfdata("pending", len(pending), "", args.warning, args.critical, 0, len(data)),
        perfdata("failed", len(failed), "", None, None, 0, len(data)),
    ]
    if failed:
        return CRITICAL, "%d batched background migration(s) failed: %s" % (len(failed), ", ".join(failed)), perf

    state = evaluate(args, len(pending))
    if pending:
        message = "%d of the %d most recent batched background migrations are not finished: %s" % (
            len(pending), len(data), ", ".join("%s (%s)" % (m["job_class_name"], m["status"]) for m in pending),
        )
    else:
        message = "all of the %d most recent batched background migrations are finished" % len(data)
    return state, message, perf


def rotate(args):
    path = args.token_file or os.environ.get("GITLAB_TOKEN_FILE")
    if not path:
        raise PluginError(UNKNOWN, "rotation writes the new token back, give the token file with -f or "
                                   "$GITLAB_TOKEN_FILE")
    if args.lifetime <= args.rotate_before:
        raise PluginError(UNKNOWN, "--lifetime must be larger than --rotate-before, or every run rotates")

    with TokenFile(path) as token_file:
        token = token_file.read()

        # Read first: a revoked token must never reach the rotate endpoint, GitLab revokes its whole family.
        data, _ = api_get(args, "/personal_access_tokens/self", token)
        name = data.get("name", "?")
        expires = data.get("expires_at")
        if not expires:
            return OK, "access token %r never expires, not rotated" % name, []

        today = datetime.datetime.now(datetime.timezone.utc).date()
        days = (datetime.date.fromisoformat(expires[:10]) - today).days
        if days >= args.rotate_before:
            return OK, "access token %r is valid for %d more days, rotation is due below %d" % (
                name, days, args.rotate_before), []

        new_expiry = (today + datetime.timedelta(days=args.lifetime)).isoformat()
        token_file.prepare()
        try:
            answer, _ = http_request(args, "/personal_access_tokens/self/rotate", token, "POST",
                                     {"expires_at": new_expiry})
            new_token = answer["token"]
        except urllib.error.HTTPError as err:
            if err.code < 500:
                raise PluginError(UNKNOWN, "GitLab refused the rotation with HTTP %d, the stored token is "
                                           "unchanged and still valid" % err.code)
            raise PluginError(CRITICAL, uncertain_rotation("HTTP %d" % err.code))
        except (OSError, ValueError, KeyError, TypeError) as err:
            raise PluginError(CRITICAL, uncertain_rotation(err))

        try:
            token_file.commit(new_token)
        except OSError as err:
            raise PluginError(CRITICAL, "GitLab rotated the token, but storing the new one failed (%s); the old "
                                        "token is revoked, create a new one" % err)

    try:
        api_get(args, "/personal_access_tokens/self", new_token)
    except PluginError as err:
        raise PluginError(CRITICAL, "stored the rotated token, but it does not work: %s" % err.message)

    return OK, "rotated access token %r, the new one is valid until %s" % (name, new_expiry), []


def uncertain_rotation(reason):
    return ("rotation request failed (%s) and GitLab may have processed it anyway; the stored token is "
            "unchanged, if it is rejected now, create a new token" % reason)


MODES = {
    "health": (check_health, "check that the API answers and report the version", None, None),
    "token": (check_token, "check how many days the access token is still valid", "30:", "7:"),
    "stats": (check_stats, "report users, projects, groups, issues and merge requests", None, None),
    "sidekiq": (check_sidekiq, "check the Sidekiq background job processing", None, None),
    "migrations": (check_migrations, "check the batched background migrations", None, None),
}


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_gitlab",
        description="Nagios/Icinga plugin to check a GitLab instance via its REST API.",
    )
    parser.add_argument("-V", "--version", action="version", version="check_gitlab %s" % VERSION)

    connection = ArgumentParser(add_help=False)
    connection.add_argument("-H", "--hostname", required=True, help="host name or address of the GitLab server")
    connection.add_argument("-p", "--port", type=int, default=DEFAULT_PORT,
                            help="port of the GitLab server (default: %d)" % DEFAULT_PORT)
    connection.add_argument("-S", "--ssl", dest="ssl", action="store_true", default=True,
                            help="use HTTPS (default)")
    connection.add_argument("--no-ssl", dest="ssl", action="store_false", help="use plain HTTP")
    connection.add_argument("-k", "--insecure", action="store_true", help="do not verify the TLS certificate")
    connection.add_argument("-t", "--timeout", type=float, default=DEFAULT_TIMEOUT,
                            help="timeout in seconds (default: %d)" % DEFAULT_TIMEOUT)

    check = ArgumentParser(add_help=False)
    check.add_argument("-T", "--token", help="access token; visible in the process list, prefer -f")
    check.add_argument("-f", "--token-file", help="file holding the access token")
    check.add_argument("-w", "--warning", help="warning threshold as a range")
    check.add_argument("-c", "--critical", help="critical threshold as a range")

    subparsers = parser.add_subparsers(dest="mode", required=True, metavar="MODE")
    for name, (_, description, warning, critical) in sorted(MODES.items()):
        defaults = []
        if warning:
            defaults.append("-w %s" % warning)
        if critical:
            defaults.append("-c %s" % critical)
        help_text = description
        if defaults:
            help_text += " (defaults to %s)" % " ".join(defaults)
        subparsers.add_parser(name, parents=[connection, check], help=help_text, description=help_text)

    help_text = "rotate the access token once it gets close to expiring and store the new one"
    rotation = subparsers.add_parser("rotate", parents=[connection], help=help_text, description=help_text)
    rotation.add_argument("-f", "--token-file", help="file holding the access token, the new token replaces it")
    rotation.add_argument("--lifetime", type=int, default=DEFAULT_LIFETIME,
                          help="days the new token is valid (default: %d)" % DEFAULT_LIFETIME)
    rotation.add_argument("--rotate-before", type=int, default=DEFAULT_ROTATE_BEFORE,
                          help="rotate once fewer days are left (default: %d)" % DEFAULT_ROTATE_BEFORE)

    return parser


def main():
    args = build_parser().parse_args()

    try:
        if args.mode == "rotate":
            state, message, perf = rotate(args)
        else:
            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
            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("GITLAB %s UNKNOWN - unexpected answer from the server: %r" % (args.mode.upper(), err))
        return UNKNOWN
    except PluginError as err:
        print("GITLAB %s %s - %s" % (args.mode.upper(), STATE_TEXT[err.state], err.message))
        return err.state

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