#!/usr/bin/python3
# Nagios/Icinga plugin to check an Immich 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 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.3"
DEFAULT_PORT = 2283
DEFAULT_TIMEOUT = 10


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, no size suffixes."""
        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, parse_value=float):
        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_value(low)
            end = float("inf") if high == "" else parse_value(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 parse_size(text):
    """Parse a byte size, optionally suffixed with K, M, G, T or P."""
    factors = {"K": 1024, "M": 1024 ** 2, "G": 1024 ** 3, "T": 1024 ** 4, "P": 1024 ** 5}
    text = text.strip()
    if text and text[-1].upper() in factors:
        return float(text[:-1]) * factors[text[-1].upper()]
    return float(text)


def human_bytes(value):
    value = float(value)
    for unit in ("B", "KiB", "MiB", "GiB", "TiB"):
        if abs(value) < 1024 or unit == "TiB":
            return "%d %s" % (value, unit) if unit == "B" else "%.1f %s" % (value, unit)
        value /= 1024


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

    path = args.token_file or os.environ.get("IMMICH_API_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("IMMICH_API_TOKEN")
    if token:
        return token.strip()

    return None


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

    request = urllib.request.Request(url, headers={"Accept": "application/json"}, method=method)
    if token:
        request.add_header("x-api-key", token)

    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%s" % ("https" if args.ssl else "http", args.hostname, args.port, path)
    try:
        return http_request(args, path, token)
    except urllib.error.HTTPError as err:
        if err.code in (401, 403):
            raise PluginError(UNKNOWN, "API token rejected for %s (HTTP %d)" % (path, err.code))
        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" % path)


def require_token(args):
    token = read_token(args)
    if not token:
        raise PluginError(UNKNOWN, "no API token given, use -T, -f or $IMMICH_API_TOKEN")
    return token


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):
    answer, elapsed = api_get(args, "/server/ping")
    if answer.get("res") != "pong":
        raise PluginError(CRITICAL, "server did not answer the ping, got %r" % answer)

    detail = ""
    try:
        version, version_elapsed = api_get(args, "/server/version")
        elapsed += version_elapsed
        detail = " (version %d.%d.%d)" % (
            version.get("major", 0),
            version.get("minor", 0),
            version.get("patch", 0),
        )
    except PluginError:
        pass

    state = OK
    if args.critical and args.critical.breached(elapsed):
        state = CRITICAL
    elif args.warning and args.warning.breached(elapsed):
        state = WARNING

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


def check_storage(args):
    token = require_token(args)
    data, _ = api_get(args, "/server/storage", token)

    try:
        usage = float(data["diskUsagePercentage"])
        used = int(data["diskUseRaw"])
        total = int(data["diskSizeRaw"])
        available = int(data["diskAvailableRaw"])
    except (KeyError, TypeError, ValueError):
        raise PluginError(UNKNOWN, "unexpected answer from /server/storage")

    state = OK
    if args.critical and args.critical.breached(usage):
        state = CRITICAL
    elif args.warning and args.warning.breached(usage):
        state = WARNING

    message = "%.2f%% of the storage used (%s of %s, %s available)" % (
        usage,
        human_bytes(used),
        human_bytes(total),
        human_bytes(available),
    )
    perf = [
        perfdata("usage", round(usage, 2), "%", args.warning, args.critical, 0, 100),
        perfdata("used", used, "B", None, None, 0, total),
        perfdata("available", available, "B", None, None, 0, total),
    ]
    return state, message, perf


def check_stats(args):
    token = require_token(args)
    data, _ = api_get(args, "/server/statistics", token)

    try:
        photos = int(data["photos"])
        videos = int(data["videos"])
        usage = int(data["usage"])
        usage_photos = int(data.get("usagePhotos", 0))
        usage_videos = int(data.get("usageVideos", 0))
    except (KeyError, TypeError, ValueError):
        raise PluginError(UNKNOWN, "unexpected answer from /server/statistics")

    state = OK
    if args.critical and args.critical.breached(usage):
        state = CRITICAL
    elif args.warning and args.warning.breached(usage):
        state = WARNING

    users = data.get("usageByUser") or []
    message = "%d photos, %d videos, %s used by %d user(s)" % (
        photos,
        videos,
        human_bytes(usage),
        len(users),
    )
    perf = [
        perfdata("photos", photos, "", None, None, 0),
        perfdata("videos", videos, "", None, None, 0),
        perfdata("usage", usage, "B", args.warning, args.critical, 0),
        perfdata("usage_photos", usage_photos, "B", None, None, 0),
        perfdata("usage_videos", usage_videos, "B", None, None, 0),
        perfdata("users", len(users), "", None, None, 0),
    ]
    return state, message, perf


def rotate(args):
    path = args.token_file or os.environ.get("IMMICH_API_TOKEN_FILE")
    if not path:
        raise PluginError(UNKNOWN, "rotation writes the new key back, give the token file with -f or "
                                   "$IMMICH_API_TOKEN_FILE")

    with TokenFile(path) as token_file:
        token = token_file.read()
        key, _ = api_get(args, "/api-keys/me", token)
        name = key.get("name", "?")

        token_file.prepare()
        try:
            answer, _ = http_request(args, "/api-keys/%s/rotate" % key["id"], token, "POST")
            new_token = answer["secret"]
        except urllib.error.HTTPError as err:
            if err.code < 500:
                raise PluginError(UNKNOWN, "Immich refused the rotation with HTTP %d, the stored key is unchanged "
                                           "and still valid; the key needs the apiKey.rotate permission" % 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, "Immich rotated the key, but storing the new one failed (%s); the old "
                                        "key is invalid, create a new one" % err)

    try:
        api_get(args, "/api-keys/me", new_token)
    except PluginError as err:
        raise PluginError(CRITICAL, "stored the rotated key, but it does not work: %s" % err.message)

    return OK, "rotated API key %r" % name, []


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


MODES = {
    "health": (check_health, "check that the server answers its ping endpoint", float),
    "storage": (check_storage, "check the storage usage of the server", float),
    "stats": (check_stats, "report the asset statistics of the server", parse_size),
}


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

    connection = ArgumentParser(add_help=False)
    connection.add_argument("-H", "--hostname", required=True, help="host name or address of the Immich server")
    connection.add_argument("-p", "--port", type=int, default=DEFAULT_PORT,
                            help="port of the Immich server (default: %d)" % DEFAULT_PORT)
    connection.add_argument("-S", "--ssl", action="store_true", help="use HTTPS instead of 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="API token; visible in the process list, prefer -f")
    check.add_argument("-f", "--token-file", help="file holding the API 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, _) in sorted(MODES.items()):
        subparsers.add_parser(name, parents=[connection, check], help=description, description=description)

    help_text = "rotate the API key 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 API token, the new token replaces it")

    return parser


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

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

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