#!/usr/bin/python3
# -*- coding: utf-8 -*-
#
# Part-DB monitoring plugin for Nagios/Naemon/Icinga
# Copyright (C) 2026 Thomas Wagner
#
# SPDX-License-Identifier: GPL-2.0-only
#
# This program is free software; you can redistribute it and/or modify it
# under the terms of the GNU General Public License version 2 as published
# by the Free Software Foundation. It is distributed in the hope that it
# will be useful, but WITHOUT ANY WARRANTY; without even the implied
# warranty of MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
# LICENSE file for details.

"""Nagios/Naemon/Icinga plugin for Part-DB, Python 3 standard library only."""

import argparse
import datetime
import json
import os
import re
import socket
import ssl
import sys
import time
import urllib.error
import urllib.parse
import urllib.request
from concurrent.futures import ThreadPoolExecutor

VERSION = "0.2"

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

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

# Order in which states escalate: a plugin's overall state is the "worst" one.
_SEVERITY = {OK: 0, WARNING: 1, UNKNOWN: 2, CRITICAL: 3}

TOKEN_LEVELS = {1: "read-only", 2: "edit", 3: "admin", 4: "full"}


class PluginError(Exception):
    """Usage or protocol problem; the plugin exits UNKNOWN with this message."""


class AuthFailure(PluginError):
    """The API rejected the credentials: monitoring lost its access, so this is CRITICAL."""


class ApiUnreachable(PluginError):
    """The instance could not be contacted at all, or TLS failed: a service failure, CRITICAL."""


# --------------------------------------------------------------------------
# Threshold ranges (monitoring-plugins guideline format)
#   10      alert if value < 0 or value > 10
#   10:     alert if value < 10
#   ~:10    alert if value > 10
#   10:20   alert if value < 10 or value > 20
#   @10:20  alert if 10 <= value <= 20 (inverted)
# --------------------------------------------------------------------------
class Range(object):
    def __init__(self, spec):
        self.spec = spec
        self.inverted = False
        self.start = 0.0
        self.end = float("inf")

        text = str(spec).strip()
        if not text:
            raise PluginError("empty threshold range")
        if text.startswith("@"):
            self.inverted = True
            text = text[1:]

        if ":" in text:
            low, high = text.split(":", 1)
            self.start = float("-inf") if low in ("~", "") else self._num(low)
            self.end = float("inf") if high == "" else self._num(high)
        else:
            self.end = self._num(text)

        if self.start > self.end:
            raise PluginError("invalid threshold range '%s': start is greater than end" % spec)

    @staticmethod
    def _num(text):
        try:
            return float(text)
        except ValueError:
            raise PluginError("invalid number '%s' in threshold range" % text)

    def breached(self, value):
        """True when *value* should raise an alert for this range."""
        inside = self.start <= value <= self.end
        return inside if self.inverted else not inside

    def __str__(self):
        return self.spec


def parse_range(spec):
    return None if spec is None else Range(spec)


# --------------------------------------------------------------------------
# Performance data and result accumulation
# --------------------------------------------------------------------------
class Perfdata(object):
    def __init__(self, label, value, uom="", warn=None, crit=None, minimum=None, maximum=None):
        self.label = label
        self.value = value
        self.uom = uom
        self.warn = warn
        self.crit = crit
        self.minimum = minimum
        self.maximum = maximum

    @staticmethod
    def _fmt(value):
        if value is None:
            return ""
        if isinstance(value, float):
            if value != value or value in (float("inf"), float("-inf")):
                return ""
            text = "%.6f" % value
            text = text.rstrip("0").rstrip(".")
            return text if text else "0"
        return str(value)

    def __str__(self):
        # Labels are always single quoted: valid everywhere, and safe for spaces or '='.
        label = str(self.label).replace("'", "")
        parts = [
            "'%s'=%s%s" % (label, self._fmt(self.value), self.uom),
            self._fmt(self.warn),
            self._fmt(self.crit),
            self._fmt(self.minimum),
            self._fmt(self.maximum),
        ]
        while len(parts) > 1 and parts[-1] == "":
            parts.pop()
        return ";".join(parts)


class Plugin(object):
    def __init__(self, shortname):
        self.shortname = shortname
        self.status = OK
        self.messages = {OK: [], WARNING: [], CRITICAL: [], UNKNOWN: []}
        self.perfdata = []
        self.extra_lines = []

    def add_status(self, code, message=None):
        if _SEVERITY[code] > _SEVERITY[self.status]:
            self.status = code
        if message:
            self.messages[code].append(message)

    def add_perfdata(self, *args, **kwargs):
        self.perfdata.append(Perfdata(*args, **kwargs))

    def add_line(self, text):
        """Extra long-output line, shown below the summary."""
        self.extra_lines.append(text)

    def check_value(self, value, warn=None, crit=None):
        """Return the state *value* falls into for the given ranges."""
        if crit is not None and crit.breached(value):
            return CRITICAL
        if warn is not None and warn.breached(value):
            return WARNING
        return OK

    def check_and_report(self, value, warn, crit, template):
        """Evaluate thresholds and record a message built from *template* ({value})."""
        code = self.check_value(value, warn, crit)
        self.add_status(code, template.format(value=value))
        return code

    def exit(self, summary=None):
        pieces = []
        for code in (CRITICAL, UNKNOWN, WARNING, OK):
            pieces.extend(self.messages[code])
        text = ", ".join(pieces) if pieces else (summary or "")

        line = "%s %s - %s" % (self.shortname, STATUS_TEXT[self.status], text)
        if self.perfdata:
            line += " | " + " ".join(str(p) for p in self.perfdata)
        print(line)
        for extra in self.extra_lines:
            print(extra)
        sys.exit(self.status)


# --------------------------------------------------------------------------
# HTTP / API access
# --------------------------------------------------------------------------
class Response(object):
    def __init__(self, status, headers, body, elapsed, url):
        self.status = status
        self.headers = headers
        self.body = body
        self.elapsed = elapsed
        self.url = url

    def json(self):
        try:
            return json.loads(self.body.decode("utf-8", "replace"))
        except ValueError as exc:
            raise PluginError("response from %s is not valid JSON: %s" % (self.url, exc))

    def text(self):
        return self.body.decode("utf-8", "replace")


def ssl_context(insecure=False, ca_cert=None):
    if insecure:
        ctx = ssl.create_default_context()
        ctx.check_hostname = False
        ctx.verify_mode = ssl.CERT_NONE
        return ctx
    return ssl.create_default_context(cafile=ca_cert)


class _NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, req, fp, code, msg, headers, newurl):
        return None


def http_get(url, token=None, timeout=10, insecure=False, ca_cert=None,
             accept="application/ld+json", follow_redirects=True):
    """GET *url*; HTTP error codes are returned, only transport failures raise."""
    request = urllib.request.Request(url, method="GET")
    request.add_header("Accept", accept)
    request.add_header("User-Agent", "check_part-db (monitoring plugin)")
    if token:
        request.add_header("Authorization", "Bearer " + token)

    handlers = [urllib.request.HTTPSHandler(context=ssl_context(insecure, ca_cert))]
    if not follow_redirects:
        handlers.append(_NoRedirect())
    opener = urllib.request.build_opener(*handlers)

    started = time.time()
    try:
        with opener.open(request, timeout=timeout) as response:
            body = response.read()
            return Response(response.status, dict(response.headers), body, time.time() - started, response.url)
    except urllib.error.HTTPError as exc:
        body = exc.read() if hasattr(exc, "read") else b""
        return Response(exc.code, dict(exc.headers or {}), body, time.time() - started, url)
    except urllib.error.URLError as exc:
        reason = exc.reason
        if isinstance(reason, ssl.SSLError):
            raise ApiUnreachable("TLS error for %s: %s (use --insecure or --ca-cert to adjust)" % (url, reason))
        if isinstance(reason, socket.timeout):
            raise ApiUnreachable("timeout after %ss connecting to %s" % (timeout, url))
        raise ApiUnreachable("cannot reach %s: %s" % (url, reason))
    except socket.timeout:
        raise ApiUnreachable("timeout after %ss reading from %s" % (timeout, url))
    except OSError as exc:
        raise ApiUnreachable("cannot reach %s: %s" % (url, exc))


def total_items(payload, url=""):
    """Read the collection size out of an API Platform / Hydra response."""
    if isinstance(payload, list):
        return len(payload)
    if isinstance(payload, dict):
        for key in ("hydra:totalItems", "totalItems"):
            if key in payload:
                return int(payload[key])
        for key in ("hydra:member", "member"):
            if key in payload:
                return len(payload[key])
    raise PluginError("cannot determine item count from response%s" % (" of " + url if url else ""))


def collection_members(payload):
    """Return the item list out of an API Platform / Hydra response."""
    if isinstance(payload, list):
        return payload
    if isinstance(payload, dict):
        for key in ("hydra:member", "member"):
            if key in payload:
                return payload[key]
    return []


def _short(text, limit=160):
    text = " ".join(text.split())
    return text if len(text) <= limit else text[:limit] + "..."


class PartDbApi(object):
    def __init__(self, base_url, token, timeout=10, insecure=False, ca_cert=None):
        self.base_url = base_url.rstrip("/")
        self.token = token
        self.timeout = timeout
        self.insecure = insecure
        self.ca_cert = ca_cert

    def url_for(self, path, params=None):
        if path.startswith("http://") or path.startswith("https://"):
            url = path
        else:
            url = self.base_url + "/" + path.lstrip("/")
        if params:
            url += ("&" if "?" in url else "?") + urllib.parse.urlencode(params)
        return url

    def get(self, path, params=None, accept="application/ld+json"):
        return http_get(self.url_for(path, params), token=self.token, timeout=self.timeout,
                        insecure=self.insecure, ca_cert=self.ca_cert, accept=accept)

    def get_json(self, path, params=None, accept="application/ld+json"):
        """GET and decode, turning API level failures into PluginError."""
        response = self.get(path, params, accept)
        if response.status == 401:
            raise AuthFailure("API token rejected (HTTP 401) - check the token and that it has not expired")
        if response.status == 403:
            raise AuthFailure("API token lacks permission for %s (HTTP 403)" % response.url)
        if response.status >= 400:
            raise PluginError("HTTP %s from %s: %s" % (response.status, response.url, _short(response.text())))
        return response.json()


# --------------------------------------------------------------------------
# Command line handling
# --------------------------------------------------------------------------
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 resolve_token(args):
    """Token from --token, --token-file, $PART_DB_TOKEN_FILE or $PART_DB_TOKEN."""
    if getattr(args, "token", None):
        return args.token.strip()

    path = getattr(args, "token_file", None) or os.environ.get("PART_DB_TOKEN_FILE")
    if path:
        try:
            with open(os.path.expanduser(path), "r") as handle:
                token = handle.read().strip()
        except IOError as exc:
            raise PluginError("cannot read token file %s: %s" % (path, exc))
        if not token:
            raise PluginError("token file %s is empty" % path)
        return token

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

    raise PluginError("no API token given - use --token, --token-file, $PART_DB_TOKEN or $PART_DB_TOKEN_FILE")


def base_url_from_args(args):
    """Assemble the base URL from -H/-p/-S/-u."""
    host = args.hostname.strip()
    if "://" in host or "/" in host:
        parts = urllib.parse.urlsplit(host if "://" in host else "//" + host)
        hint = "-H %s" % (parts.hostname or "HOST")
        try:
            if parts.port:
                hint += " -p %d" % parts.port
        except ValueError:
            pass
        if parts.path.strip("/"):
            hint += " -u %s" % parts.path.rstrip("/")
        if parts.scheme == "http":
            hint += " --no-ssl"
        raise PluginError("-H takes a host name or address, not a URL; use %s" % hint)
    if not host:
        raise PluginError("-H requires a host name or address")

    scheme = "https" if args.ssl else "http"
    port = args.port if args.port else (443 if args.ssl else 80)
    if port < 1 or port > 65535:
        raise PluginError("invalid port %s" % port)

    # Bracket IPv6 literals so the URL stays parseable.
    if ":" in host and not host.startswith("["):
        host = "[%s]" % host
    # Leave the default port out: it keeps the Host header conventional.
    authority = host if port == (443 if args.ssl else 80) else "%s:%d" % (host, port)

    uri = args.uri.strip()
    if uri and not uri.startswith("/"):
        uri = "/" + uri
    return "%s://%s%s" % (scheme, authority, uri.rstrip("/"))


def api_from_args(args):
    return PartDbApi(base_url_from_args(args), resolve_token(args),
                     timeout=args.timeout, insecure=args.insecure, ca_cert=args.ca_cert)


# --------------------------------------------------------------------------
# health: the application answers, the token is accepted, the database answers
# --------------------------------------------------------------------------
def mode_health(args, plugin):
    warn = parse_range(args.warning)
    crit = parse_range(args.critical)
    api = api_from_args(args)

    entry = api.get("/api")
    if entry.status == 401:
        plugin.add_status(CRITICAL, "API token rejected (HTTP 401)")
        plugin.add_perfdata("api_time", round(entry.elapsed, 4), "s")
        plugin.exit()
    if entry.status == 403:
        plugin.add_status(CRITICAL, "API token has no access to the entrypoint (HTTP 403)")
        plugin.exit()
    if entry.status >= 500:
        plugin.add_status(CRITICAL, "Part-DB returned HTTP %s on /api (application or backend error)" % entry.status)
        plugin.exit()
    if entry.status != 200:
        plugin.add_status(CRITICAL, "unexpected HTTP %s on /api" % entry.status)
        plugin.exit()

    try:
        collections = [key for key in entry.json() if not key.startswith("@")]
    except PluginError:
        plugin.add_status(CRITICAL, "/api did not return a usable API entrypoint document")
        plugin.exit()

    entry_time = entry.elapsed

    # A real query, so a broken or unreachable database is caught.
    probe = api.get(args.probe_path, {"itemsPerPage": 1, "page": 1})
    if probe.status >= 500:
        plugin.add_status(CRITICAL, "database query %s failed with HTTP %s" % (args.probe_path, probe.status))
        plugin.add_perfdata("api_time", round(entry_time, 4), "s")
        plugin.exit()
    if probe.status != 200:
        plugin.add_status(CRITICAL, "query %s returned HTTP %s" % (args.probe_path, probe.status))
        plugin.exit()

    try:
        items = total_items(probe.json(), args.probe_path)
    except PluginError as exc:
        plugin.add_status(CRITICAL, str(exc))
        plugin.exit()

    total_time = entry_time + probe.elapsed

    plugin.check_and_report(
        round(total_time, 3), warn, crit,
        "API responding in {value}s, token valid, %d collections, %d objects in %s"
        % (len(collections), items, args.probe_path))

    plugin.add_perfdata("time", round(total_time, 4), "s", warn=warn, crit=crit, minimum=0)
    plugin.add_perfdata("api_time", round(entry_time, 4), "s", minimum=0)
    plugin.add_perfdata("query_time", round(probe.elapsed, 4), "s", minimum=0)
    plugin.add_perfdata("collections", len(collections), minimum=0)


# --------------------------------------------------------------------------
# web: the frontend page and its TLS certificate, no token needed
# --------------------------------------------------------------------------
VERSION_RE = re.compile(r"Version:\s*([0-9][0-9A-Za-z.\-+]*)")
GIT_RE = re.compile(r"Version:\s*[0-9][0-9A-Za-z.\-+]*\s*\(/([0-9a-f]+)\)")


def certificate_days_left(url, timeout, ca_cert=None):
    """Days until the TLS certificate of *url* expires.

    Python hands out an empty certificate when verification is off, which is
    why --insecure skips this check instead of reporting on it."""
    parts = urllib.parse.urlsplit(url)
    host = parts.hostname
    port = parts.port or 443

    context = ssl.create_default_context(cafile=ca_cert)
    try:
        with socket.create_connection((host, port), timeout=timeout) as raw:
            with context.wrap_socket(raw, server_hostname=host) as tls:
                cert = tls.getpeercert()
    except ssl.SSLCertVerificationError as exc:
        raise PluginError("certificate not trusted: %s" % exc.verify_message)
    except ssl.SSLError as exc:
        raise PluginError("TLS handshake failed: %s" % exc)
    except (socket.timeout, OSError) as exc:
        raise ApiUnreachable("cannot open a TLS connection to %s:%s: %s" % (host, port, exc))

    if not cert or "notAfter" not in cert:
        raise PluginError("peer presented no usable certificate")

    expires = datetime.datetime.fromtimestamp(ssl.cert_time_to_seconds(cert["notAfter"]), datetime.timezone.utc)
    delta = expires - datetime.datetime.now(datetime.timezone.utc)
    return delta.total_seconds() / 86400.0, expires


def mode_web(args, plugin):
    warn = parse_range(args.warning)
    crit = parse_range(args.critical)

    base = base_url_from_args(args) + "/"
    response = http_get(base, timeout=args.timeout, insecure=args.insecure, ca_cert=args.ca_cert, accept="text/html")

    if response.status >= 500:
        plugin.add_status(CRITICAL, "frontend returned HTTP %s" % response.status)
        plugin.exit()
    elif response.status != 200:
        plugin.add_status(CRITICAL, "frontend returned HTTP %s (expected 200)" % response.status)
        plugin.exit()

    body = response.text()

    if args.expect_string and args.expect_string not in body:
        plugin.add_status(CRITICAL, "page does not contain '%s'" % args.expect_string)

    version_match = VERSION_RE.search(body)
    version = version_match.group(1) if version_match else None
    git_match = GIT_RE.search(body)

    described = "Part-DB %s" % version if version else "frontend"
    if git_match:
        described += " (%s)" % git_match.group(1)

    plugin.check_and_report(round(response.elapsed, 3), warn, crit,
                            "%s served in {value}s, %d bytes" % (described, len(response.body)))

    plugin.add_perfdata("time", round(response.elapsed, 4), "s", warn=warn, crit=crit, minimum=0)
    plugin.add_perfdata("size", len(response.body), "B", minimum=0)

    if args.no_cert_check or not base.lower().startswith("https://"):
        pass
    elif args.insecure:
        plugin.add_line("Certificate expiry not checked: --insecure turns certificate validation off, "
                        "and an unvalidated certificate cannot be read.")
    else:
        try:
            days, expires = certificate_days_left(base, args.timeout, args.ca_cert)
        except PluginError as exc:
            # Do not mask a working page: report the certificate problem next to the page result.
            plugin.add_status(CRITICAL, str(exc))
        else:
            if days < args.cert_critical:
                plugin.add_status(CRITICAL, "certificate expires in %.1f days" % days)
            elif days < args.cert_warning:
                plugin.add_status(WARNING, "certificate expires in %.1f days" % days)
            else:
                plugin.add_status(OK, "certificate valid for %.0f days" % days)
            plugin.add_perfdata("cert_days", round(days, 2), "", warn=args.cert_warning, crit=args.cert_critical)
            plugin.add_line("Certificate expires %s UTC" % expires.strftime("%Y-%m-%d %H:%M:%S"))

    if response.url != base:
        plugin.add_line("Final URL after redirects: %s" % response.url)


# --------------------------------------------------------------------------
# stats: object counts of every collection the API offers
# --------------------------------------------------------------------------
SUMMARY_ENTITIES = ("parts", "categories", "footprints", "manufacturers",
                    "suppliers", "storage_locations", "part_lots", "projects")


def parse_thresholds(specs):
    """--threshold parts=900:,500: -> {'parts': (Range, Range)}"""
    thresholds = {}
    for spec in specs:
        if "=" not in spec:
            raise PluginError("invalid --threshold '%s', expected ENTITY=WARNING[,CRITICAL]" % spec)
        name, ranges = spec.split("=", 1)
        pieces = ranges.split(",")
        if len(pieces) > 2:
            raise PluginError("invalid --threshold '%s', expected at most one warning and one critical range" % spec)
        warn = parse_range(pieces[0]) if pieces[0] else None
        crit = parse_range(pieces[1]) if len(pieces) == 2 and pieces[1] else None
        thresholds[name.strip()] = (warn, crit)
    return thresholds


def discover(api):
    """Map collection name -> path, from the API entrypoint."""
    payload = api.get_json("/api")
    if not isinstance(payload, dict):
        raise PluginError("/api did not return an API entrypoint document")
    collections = {}
    for key, value in payload.items():
        if key.startswith("@") or not isinstance(value, str):
            continue
        # The path's last segment ('storage_locations') makes a steadier perfdata label than the camelCase key.
        collections[value.rstrip("/").rsplit("/", 1)[-1]] = value
    if not collections:
        raise PluginError("API entrypoint listed no collections")
    return collections


def count_one(api, name, path):
    """Return (name, count, error_message)."""
    try:
        response = api.get(path, {"itemsPerPage": 1, "page": 1})
    except PluginError as exc:
        return name, None, str(exc)
    if response.status == 403:
        return name, None, "not readable with this token (HTTP 403)"
    if response.status != 200:
        return name, None, "HTTP %s" % response.status
    try:
        return name, total_items(response.json(), path), None
    except PluginError as exc:
        return name, None, str(exc)


def mode_stats(args, plugin):
    thresholds = parse_thresholds(args.threshold)
    api = api_from_args(args)
    available = discover(api)

    if args.list:
        plugin.add_status(OK, "%d collections available" % len(available))
        for name in sorted(available):
            plugin.add_line("%-24s %s" % (name, available[name]))
        return

    if args.entity:
        selected = {}
        for name in args.entity:
            if name not in available:
                raise PluginError("unknown entity '%s'; available: %s" % (name, ", ".join(sorted(available))))
            selected[name] = available[name]
    else:
        selected = available

    unknown_thresholds = set(thresholds) - set(selected)
    if unknown_thresholds:
        raise PluginError("--threshold given for entity not being counted: %s" % ", ".join(sorted(unknown_thresholds)))

    workers = max(1, min(args.workers, len(selected)))
    with ThreadPoolExecutor(max_workers=workers) as pool:
        results = list(pool.map(lambda item: count_one(api, item[0], item[1]), sorted(selected.items())))

    counts = {}
    skipped = []
    for name, count, error in results:
        if error is not None:
            skipped.append((name, error))
            continue
        counts[name] = count

    if not counts:
        raise PluginError("no collection could be counted (%s)" % "; ".join("%s: %s" % item for item in skipped))

    for name in sorted(counts):
        warn, crit = thresholds.get(name, (None, None))
        code = plugin.check_value(counts[name], warn, crit)
        if code != OK:
            plugin.add_status(code, "%s=%d outside threshold" % (name, counts[name]))
        plugin.add_perfdata(name, counts[name], warn=warn, crit=crit, minimum=0)

    for name, error in skipped:
        if args.strict:
            plugin.add_status(UNKNOWN, "%s: %s" % (name, error))
        else:
            plugin.add_line("skipped %s: %s" % (name, error))

    highlights = [(name, counts[name]) for name in SUMMARY_ENTITIES if name in counts]
    if not highlights:
        highlights = sorted(counts.items(), key=lambda kv: -kv[1])[:4]
    summary = ", ".join("%d %s" % (value, name) for name, value in highlights)
    if skipped and not args.strict:
        summary += " (%d collection(s) skipped)" % len(skipped)
    plugin.add_status(OK, summary)

    for name in sorted(counts):
        plugin.add_line("%-24s %d" % (name, counts[name]))


# --------------------------------------------------------------------------
# stock: parts below their configured minimum amount
#
# Part-DB computes a part's stock from its lots, and that computed field
# cannot be filtered on server side, so only parts with minamount > 0 are
# fetched and compared locally.
# --------------------------------------------------------------------------
def fetch_parts(api, params, page_size, max_pages):
    """Yield part objects, following pagination."""
    page = 1
    seen = 0
    while page <= max_pages:
        query = dict(params)
        query["itemsPerPage"] = page_size
        query["page"] = page
        payload = api.get_json("/api/parts", query)
        members = collection_members(payload)
        if not members:
            return
        for member in members:
            yield member
        seen += len(members)
        try:
            total = total_items(payload)
        except PluginError:
            total = None
        if total is not None and seen >= total:
            return
        if len(members) < page_size:
            return
        page += 1
    raise PluginError("stopped after %d pages; raise --max-pages or --page-size" % max_pages)


def amount(value):
    """Part-DB amounts can be fractional depending on the measurement unit."""
    try:
        return float(value)
    except (TypeError, ValueError):
        return 0.0


def fmt_amount(value):
    return "%d" % value if float(value).is_integer() else "%.3g" % value


def mode_stock(args, plugin):
    warn = parse_range(args.warning) if args.warning else None
    crit = parse_range(args.critical) if args.critical else None
    zero_warn = parse_range(args.zero_warning) if args.zero_warning else None
    zero_crit = parse_range(args.zero_critical) if args.zero_critical else None

    if (zero_warn or zero_crit) and not args.full_scan:
        raise PluginError("--zero-warning/--zero-critical require --full-scan")
    if args.page_size < 1:
        raise PluginError("--page-size must be at least 1")

    api = api_from_args(args)

    monitored = 0
    low = []
    for part in fetch_parts(api, {"minamount[gt]": 0}, args.page_size, args.max_pages):
        monitored += 1
        minimum = amount(part.get("minamount"))
        stock = amount(part.get("total_instock"))
        if stock < minimum:
            low.append((part.get("name", "?"), part.get("id"), stock, minimum))

    low.sort(key=lambda item: (item[2] - item[3], item[0]))

    code = plugin.check_value(len(low), warn, crit)
    if low:
        plugin.add_status(code, "%d of %d monitored part(s) below minimum" % (len(low), monitored))
    elif monitored:
        plugin.add_status(code, "all %d monitored part(s) at or above minimum" % monitored)
    else:
        plugin.add_status(code, "no part defines a minimum amount, nothing to check")

    plugin.add_perfdata("low_stock", len(low), warn=warn, crit=crit, minimum=0)
    plugin.add_perfdata("monitored_parts", monitored, minimum=0)

    for name, part_id, stock, minimum in low[:args.max_list]:
        plugin.add_line("%s: %s in stock, minimum %s%s" % (
            name, fmt_amount(stock), fmt_amount(minimum),
            " (%s/en/part/%s/info)" % (api.base_url, part_id) if part_id is not None else ""))
    if len(low) > args.max_list:
        plugin.add_line("... and %d more" % (len(low) - args.max_list))

    if args.full_scan:
        total_parts = 0
        zero_stock = 0
        total_stock = 0.0
        for part in fetch_parts(api, {}, args.page_size, args.max_pages):
            total_parts += 1
            stock = amount(part.get("total_instock"))
            total_stock += stock
            if stock <= 0:
                zero_stock += 1

        zero_code = plugin.check_value(zero_stock, zero_warn, zero_crit)
        plugin.add_status(zero_code, "%d of %d part(s) have no stock" % (zero_stock, total_parts))

        plugin.add_perfdata("parts", total_parts, minimum=0)
        plugin.add_perfdata("zero_stock", zero_stock, warn=zero_warn, crit=zero_crit, minimum=0)
        plugin.add_perfdata("total_stock", round(total_stock, 3), minimum=0)


# --------------------------------------------------------------------------
# token: days until the API token itself expires
# --------------------------------------------------------------------------
def parse_timestamp(text):
    if text.endswith("Z"):
        text = text[:-1] + "+00:00"
    stamp = datetime.datetime.fromisoformat(text)
    if stamp.tzinfo is None:
        stamp = stamp.replace(tzinfo=datetime.timezone.utc)
    return stamp


def mode_token(args, plugin):
    warn = parse_range(args.warning)
    crit = parse_range(args.critical)
    token = api_from_args(args).get_json("/api/tokens/current", accept="application/json")

    name = token.get("name", "?")
    level = token.get("level")
    level = TOKEN_LEVELS.get(level, str(level).lower().replace("_", "-"))

    last_used = token.get("last_time_used")
    if last_used:
        plugin.add_line("Last used %s" % last_used)

    valid_until = token.get("valid_until")
    if not valid_until:
        plugin.add_status(OK, "API token '%s' (%s) never expires" % (name, level))
        return

    try:
        expires = parse_timestamp(valid_until)
    except ValueError:
        raise PluginError("cannot read the expiry date %r of the API token" % valid_until)
    days = (expires - datetime.datetime.now(datetime.timezone.utc)).total_seconds() / 86400.0

    plugin.check_and_report(
        round(days, 1), warn, crit,
        "API token '%s' (%s) expires on %s, in {value} days" % (name, level, expires.strftime("%Y-%m-%d")))
    plugin.add_perfdata("days_left", round(days, 2), warn=warn, crit=crit)


# --------------------------------------------------------------------------
# Command line
# --------------------------------------------------------------------------
RANGE_HELP = "see https://www.monitoring-plugins.org/doc/guidelines.html#THRESHOLDFORMAT"

EPILOGS = {
    "health": """
Checks, in order, that the API entrypoint answers (the application is up), that
the API token is accepted, and that a collection query returns (the database
answers). Thresholds are in seconds; %s

Examples:
  check_part-db health -H part-db.example.org -f /etc/naemon/part-db.token
  check_part-db health -H part-db.example.org -f TOKEN -w 1 -c 3 -t 5
  check_part-db health -H part-db.example.org -p 8443 -u /partdb -f TOKEN
""" % RANGE_HELP,
    "web": """
Checks what a user hits in the browser; no API token is needed, so it keeps
working when the API is disabled or the token expired. Certificate thresholds
are days remaining and alert below them. Pass --no-cert-check for plain HTTP.

Examples:
  check_part-db web -H part-db.example.org
  check_part-db web -H part-db.example.org --cert-warning 21 --cert-critical 7
  check_part-db web -H part-db.example.org --expect-string 'Part-DB'
  check_part-db web -H part-db.example.org --no-ssl
""",
    "stats": """
Entity names are the API collection names, e.g. parts, categories, footprints,
manufacturers, suppliers, storage_locations, part_lots, projects, attachments.
Run with --list to see what this instance offers.

Thresholds use --threshold ENTITY=WARNING[,CRITICAL] with the usual range
format, so "1000:" means "alert below 1000".

Examples:
  check_part-db stats -H part-db.example.org -f /etc/naemon/part-db.token
  check_part-db stats -H part-db.example.org -f TOKEN --threshold parts=900:,500:
  check_part-db stats -H part-db.example.org -f TOKEN -e parts -e categories
""",
    "stock": """
Thresholds apply to the NUMBER of parts below their minimum. The default "-w 0"
alerts as soon as any part is below its minimum. Only parts with a minimum
amount above zero are considered; if none sets one, the check reports OK.

Examples:
  check_part-db stock -H part-db.example.org -f /etc/naemon/part-db.token
  check_part-db stock -H part-db.example.org -f TOKEN -w 5 -c 20
  check_part-db stock -H part-db.example.org -f TOKEN --full-scan
""",
    "token": """
Part-DB API tokens expire, by default one year after they were created, and
cannot be renewed through the API. Losing the token means losing monitoring,
so this mode goes CRITICAL once fewer than 30 days are left, in time to create
a new token in the web interface. Thresholds are days remaining; add -w for an
earlier warning.

Examples:
  check_part-db token -H part-db.example.org -f /etc/naemon/part-db.token
  check_part-db token -H part-db.example.org -f TOKEN -w 60: -c 14:
""",
}


def build_parser():
    parser = ArgumentParser(
        prog="check_part-db",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        description="Nagios/Naemon/Icinga plugin for Part-DB.",
    )
    parser.add_argument("-V", "--version", action="version", version="%(prog)s " + VERSION)

    connection = ArgumentParser(add_help=False)
    connection.add_argument("-H", "--hostname", required=True, metavar="ADDRESS",
                            help="Host name or address of the Part-DB server, e.g. part-db.example.org")
    connection.add_argument("-p", "--port", type=int, metavar="PORT",
                            help="Port to connect to (default: 443 with TLS, 80 with --no-ssl)")
    connection.add_argument("-S", "--ssl", dest="ssl", action="store_true", default=True,
                            help="Use HTTPS (default). Stated explicitly for readability in service definitions")
    connection.add_argument("--no-ssl", dest="ssl", action="store_false",
                            help="Use plain HTTP. Note that the API token is then sent unencrypted")
    connection.add_argument("-u", "--uri", metavar="PATH", default="",
                            help="Path prefix when Part-DB is not served from the server root, e.g. /partdb")
    connection.add_argument("-t", "--timeout", type=float, default=10.0, metavar="SEC",
                            help="Request timeout in seconds (default: 10)")
    connection.add_argument("--insecure", action="store_true", help="Do not verify the TLS certificate")
    connection.add_argument("--ca-cert", metavar="FILE", help="CA bundle used to verify the TLS certificate")

    credentials = ArgumentParser(add_help=False)
    credentials.add_argument("-T", "--token", metavar="TOKEN",
                             help="Part-DB API token. Prefer --token-file or $PART_DB_TOKEN: arguments are visible "
                                  "in the process list")
    credentials.add_argument("-f", "--token-file", metavar="FILE",
                             help="File containing the API token (default: $PART_DB_TOKEN_FILE)")

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

    def mode(name, handler, shortname, help_text, token=True):
        sub = subparsers.add_parser(
            name, parents=[connection, credentials] if token else [connection],
            help=help_text, description=help_text, epilog=EPILOGS[name],
            formatter_class=argparse.RawDescriptionHelpFormatter)
        sub.set_defaults(handler=handler, shortname=shortname)
        return sub

    health = mode("health", mode_health, "PART-DB HEALTH",
                  "check that Part-DB is up, authenticating and answering database backed queries")
    health.add_argument("-w", "--warning", metavar="RANGE", default="2",
                        help="Warning range for the total API response time in seconds (default: 2)")
    health.add_argument("-c", "--critical", metavar="RANGE", default="5",
                        help="Critical range for the total API response time in seconds (default: 5)")
    health.add_argument("--probe-path", metavar="PATH", default="/api/parts",
                        help="Collection queried to prove the database answers (default: /api/parts)")

    web = mode("web", mode_web, "PART-DB WEB",
               "check that the web frontend serves its page and that its TLS certificate is still valid",
               token=False)
    web.add_argument("-w", "--warning", metavar="RANGE", default="3",
                     help="Warning range for page load time in seconds (default: 3)")
    web.add_argument("-c", "--critical", metavar="RANGE", default="8",
                     help="Critical range for page load time in seconds (default: 8)")
    web.add_argument("--expect-string", metavar="TEXT", default="Part-DB",
                     help="Text that must appear in the page (default: Part-DB). Empty value disables it")
    web.add_argument("--cert-warning", type=int, default=30, metavar="DAYS",
                     help="Warn when the certificate expires in fewer than this many days (default: 30)")
    web.add_argument("--cert-critical", type=int, default=14, metavar="DAYS",
                     help="Alert when the certificate expires in fewer than this many days (default: 14)")
    web.add_argument("--no-cert-check", action="store_true",
                     help="Skip the certificate check (plain HTTP, or when the certificate is monitored elsewhere)")

    stats = mode("stats", mode_stats, "PART-DB STATS",
                 "collect inventory statistics as performance data, with optional thresholds per entity")
    stats.add_argument("-e", "--entity", action="append", metavar="NAME", default=[],
                       help="Count only this entity; repeatable (default: every collection the API offers)")
    stats.add_argument("--threshold", action="append", metavar="SPEC", default=[],
                       help="ENTITY=WARNING[,CRITICAL] range spec; repeatable")
    stats.add_argument("--list", action="store_true", help="List the collections this instance offers, then exit OK")
    stats.add_argument("--strict", action="store_true",
                       help="Treat collections the token may not read as UNKNOWN instead of silently skipping them")
    stats.add_argument("--workers", type=int, default=5, metavar="N", help="Parallel API requests (default: 5)")

    stock = mode("stock", mode_stock, "PART-DB STOCK", "check for parts below their configured minimum amount")
    stock.add_argument("-w", "--warning", metavar="RANGE", default="0",
                       help="Warning range for the number of parts below their minimum (default: 0, i.e. any)")
    stock.add_argument("-c", "--critical", metavar="RANGE", default=None,
                       help="Critical range for the number of parts below their minimum (default: none)")
    stock.add_argument("--full-scan", action="store_true",
                       help="Also walk every part to report total stock and the number of parts with no stock at all")
    stock.add_argument("--zero-warning", metavar="RANGE", default=None,
                       help="Warning range for parts with zero stock; requires --full-scan")
    stock.add_argument("--zero-critical", metavar="RANGE", default=None,
                       help="Critical range for parts with zero stock; requires --full-scan")
    stock.add_argument("--page-size", type=int, default=100, metavar="N",
                       help="Parts fetched per API request (default: 100)")
    stock.add_argument("--max-pages", type=int, default=100, metavar="N",
                       help="Safety limit on pages fetched (default: 100)")
    stock.add_argument("--max-list", type=int, default=20, metavar="N",
                       help="Parts named in the long output (default: 20)")

    token = mode("token", mode_token, "PART-DB TOKEN", "check how many days the API token is still valid")
    token.add_argument("-w", "--warning", metavar="RANGE", default=None,
                       help="Warning range for the days until the token expires (default: none)")
    token.add_argument("-c", "--critical", metavar="RANGE", default="30:",
                       help="Critical range for the days until the token expires (default: 30:)")

    return parser


def main():
    args = build_parser().parse_args()
    plugin = Plugin(args.shortname)
    try:
        args.handler(args, plugin)
    except (ApiUnreachable, AuthFailure) as exc:
        print("%s CRITICAL - %s" % (args.shortname, exc))
        sys.exit(CRITICAL)
    except PluginError as exc:
        print("%s UNKNOWN - %s" % (args.shortname, exc))
        sys.exit(UNKNOWN)
    except KeyboardInterrupt:
        print("%s UNKNOWN - interrupted" % args.shortname)
        sys.exit(UNKNOWN)
    except SystemExit:
        raise
    except Exception as exc:
        print("%s UNKNOWN - unhandled %s: %s" % (args.shortname, type(exc).__name__, exc))
        sys.exit(UNKNOWN)
    plugin.exit()


if __name__ == "__main__":
    main()
