#!/usr/bin/python3

'''
Usage:

check_snmp_resources -H hostname -C community -e cpu -w 70 -c 95
check_snmp_resources -H hostname -C community -e mem -w 30 -c 50
check_snmp_resources -H hostname -C community -e disk # list partitions
check_snmp_resources -H hostname -C community -e disk -p / -w 90 -c 95
check_snmp_resources -H hostname -C community -e net -i eth0
'''

import sys
import os
import json
from netsnmp import Session, VarList, Varbind
from collections import defaultdict
from getopt import gnu_getopt

class SNMP:
  INT_TYPES = {
      "INTEGER", "INTEGER32", "UINTEGER", "UINTEGER32", "UNSIGNED32",
      "COUNTER", "COUNTER32", "COUNTER64",
      "GAUGE", "GAUGE32", "TICKS"
  }
  FLOAT_TYPES = {
      "FLOAT", "DOUBLE", "OPAQUE"
  }
  TIMEOUT = 2000000
  RETRIES = 2
  def __init__(self, host_name, community_name=None, **kwargs):
      if community_name:	# SNMP v2c
          self.session = Session(
              DestHost=host_name, Community=community_name, Version=2,
              Timeout=self.TIMEOUT, Retries=self.RETRIES
          )
      else:			# Other
          self.session = Session(
              DestHost=host_name,
              Timeout=self.TIMEOUT, Retries=self.RETRIES,
              **kwargs
          )
  def decode(self, var):
      if var.type in self.INT_TYPES:
          return int(var.val)
      elif var.type in self.FLOAT_TYPES:
          return float(var.val)
      elif var.type=="OCTETSTR" and var.tag=="ifPhysAddress":
          return ":".join(["%02x" % x for x in var.val])
      elif isinstance(var.val, bytes):
          return var.val.decode("utf-8", "replace")
      else:
          return var.val
  def check_error(self):
      if self.session.ErrorStr:
          print(f"SNMP error - {self.session.ErrorStr}")
          sys.exit(2)
  def walk(self, oid, **ids):
      rev_ids = {str(v):k for k,v in ids.items()}
      vars = VarList(*[
          Varbind(f"{oid}.{x}")
          for x in rev_ids.keys()
      ])
      self.session.walk(vars)
      self.check_error()
      return {
          rev_ids[var.iid]: self.decode(var)
          for var in vars
      }
  def get_tagged(self, oid, ids):
      vars = VarList(*[
          Varbind(f"{oid}.{x}")
          for x in ids
      ])
      self.session.get(vars)
      self.check_error()
      return {
          var.tag: self.decode(var)
          for var in vars
      }
  def get_mapped(self, oid):
      vars = VarList(Varbind(oid))
      self.session.walk(vars)
      self.check_error()
      ret = {}
      for var in vars:
          if var.iid not in ret:
              ret[var.iid] = {}
          ret[var.iid][var.tag] = self.decode(var)
      return ret

  def cpu_load(self):
      return self.walk(
          '.1.3.6.1.4.1.2021.10.1',
          load1=1,
          load5=2,
          load15=3
      )
  def memory_usage(self):
      return self.get_tagged(
          '.1.3.6.1.4.1.2021.4',
          ["6.0","5.0","15.0","4.0","3.0"]
      )
  def storage(self, fixed_disks_only=True):
      fixed_disk_type = ".1.3.6.1.2.1.25.2.1.4"
      usage = {}
      for key, data in self.get_mapped(".1.3.6.1.2.1.25.2.3.1").items():
          #print(data)
          if fixed_disks_only and data["hrStorageType"]!=fixed_disk_type:
              continue
          units = data.get("hrStorageAllocationUnits", 1024)
          usage[data["hrStorageDescr"]] = dict(
              #units = data["hrStorageAllocationUnits"],
              size = data["hrStorageSize"]*units//1024,
              used = data["hrStorageUsed"]*units//1024
          )
      return usage
  def network(self):
      ifaces = {}
      for iface in self.get_mapped(".1.3.6.1.2.1.31.1.1.1").values():
          ifname = iface["ifName"]
          ifaces[ifname] = iface
      return ifaces


def nagios_status(name, value=None, warning=None, critical=None, data=""):
    if critical is not None and value>=critical:
        print(f"{name} CRITICAL - {data}")
        sys.exit(2)
    elif warning is not None and value>=warning:
        print(f"{name} WARNING - {data}")
        sys.exit(1)
    else:
        print(f"{name} OK - {data}")
        sys.exit(0)


def getopt(opts, short, full, retype=float):
    if short in opts:
        return retype(opts[short][0])
    elif full in opts:
        return retype(opts[full][0])
    return 0


class IfaceDev:
  def __init__(self, iface):
      self.iface = iface
  def __call__(self, name, key):
      return f"{name}={self.iface[key]}c"


if __name__ == "__main__":
    snmp_opts = [
        "Version", "Community", "SecName", "SecLevel",
        "AuthProto", "AuthPass", "PrivProto", "PrivPass"
    ]
    gopts, files = gnu_getopt(sys.argv[1:], "e:H:C:J:Sw:c:p:i:", [
        "warning=", "critical=", "path=",
        *[f"{x}=" for x in snmp_opts]
    ])

    opts = defaultdict(list)
    snmp_args = {}
    for opt, arg in gopts:
        opts[opt].append(arg)
        if opt.strip("-") in snmp_opts:
            snmp_args[opt.strip("-")] = arg
    #print(opts)

    hostname = opts["-H"][0]
    if opts["-J"]:
        json_path = opts["-J"][0]
        if os.path.abspath(json_path)!=json_path:
            quit(f"Insecure json config path: {json_path}")
        elif os.path.exists(json_path):
            snmp_args = json.load(open(json_path))
        else:
            quit(f"Missing json config: {json_path}")
    if opts["-S"]:
        print(json.dumps(snmp_args, indent=2))
        quit()

    if opts["-C"]:
        snmp = SNMP(hostname, opts["-C"][0])
    else:
        snmp = SNMP(hostname, **snmp_args)
    warning = getopt(opts, "-w", "--warning")
    critical = getopt(opts, "-c", "--critical")

    if "cpu" in opts['-e']:
        status = []
        perf = []
        compare = 0
        for key, value in snmp.cpu_load().items():
            status.append(f"{value}")
            perf.append(f"{key}={value}")
            if float(value)>compare:
                compare = float(value)
        nagios_status("CPU", compare, warning, critical,
          f"load: {', '.join(status)}|{'; '.join(perf)}"
        )
    elif "mem" in opts['-e']:
        data = snmp.memory_usage()
        swap = (data["memTotalSwap"]-data["memAvailSwap"])/data["memTotalSwap"]
        nagios_status("mem", swap*100, warning, critical,
            f"{data['memTotalReal']//1024} MB mem total, {swap*100:1.0f}% swap in use"
            f"|Total={data['memTotalReal']}kiB"
            f" Used={data['memTotalReal']-data['memAvailReal']}kiB"
            f" Cached={data['memCached']}kiB"
            f" Swap={data['memTotalSwap']-data['memAvailSwap']}kiB"
        )
    elif "disk" in opts['-e'] and opts['-p']:
        part = opts['-p'][0]
        data = snmp.storage()[part]
        used = data['used']
        size = data['size']
        nagios_status("disk", used/size*100, warning, critical,
            f"{part}: {used/size*100:1.0f}% used"
            f"|{part}={used}kiB;{warning/100*size};{critical/100*size};0;{size}"
        )
    elif "disk" in opts['-e']:
        for key, value in snmp.storage().items():
            print(f"{key}: {value}")
    elif "net" in opts['-e'] and opts['-i']:
        ifname = opts['-i'][0]
        n = IfaceDev(snmp.network()[ifname])
        nagios_status("net", None, None, None,
            f"{ifname}|"
            +n('rx_bytes', 'ifHCInOctets') + ' '
            +n('tx_bytes', 'ifHCOutOctets') + ' '
            +n('rx_packets', 'ifHCInUcastPkts') + ' '
            +n('tx_packets', 'ifHCOutUcastPkts')
        )
    elif "net" in opts['-e']:
        for key, value in snmp.network().items():
            print(f"{key}")
    else:
        print("Missing parameter: -e")
        sys.exit(5)
