#!/usr/bin/python3

'''
Usage:

check_snmp_resources -H hostname -C community -e load -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
check_snmp_resources -H hostname -C community -e cisco_load
'''

import sys
import os
import json
from netsnmp import Session, VarList, Varbind
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(self, oids):
      vars = VarList(*[Varbind(x) for x in oids])
      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

  # Switches
  def cisco_load(self):
      loads = self.get([
          '.1.3.6.1.4.1.9.2.1.57.0',	# Catalyst 1min load (%)
          '.1.3.6.1.4.1.9.2.1.58.0',	# Catalyst 5min load (%)
          '.1.3.6.1.4.1.9.6.1.101.1.8.0',	# SG350 1min load (%)
          '.1.3.6.1.4.1.9.6.1.101.1.9.0'	# SG350 5min load (%)
      ])
      ret = {}
      def fractal(value):
          return value/100 if type(value)==int else None
      for oid, value in loads.items():
          value = fractal(value)
          if oid.endswith(".57.0") and value is not None:
              ret["load1"] = value
          elif oid.endswith(".58.0") and value is not None:
              ret["load5"] = value
          elif oid.endswith(".8.0") and value is not None:
              ret["load1"] = value
          elif oid.endswith(".9.0") and value is not None:
              ret["load5"] = value
      return ret

  # Linux systems
  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([
          '.1.3.6.1.4.1.2021.4.6.0',	# mem free
          '.1.3.6.1.4.1.2021.4.5.0',	# mem total
          '.1.3.6.1.4.1.2021.4.15.0',	# mem cached
          '.1.3.6.1.4.1.2021.4.4.0',	# swap free
          '.1.3.6.1.4.1.2021.4.3.0'	# swap total
      ])
  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 and value>=critical:
        print(f"{name} CRITICAL - {data}")
        sys.exit(2)
    elif warning 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])
    elif full in opts:
        return retype(opts[full])
    return None


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


def show_load(snmp_cpu_load, warning, critical):
    status = []
    perf = []
    compare = 0
    for key, value in snmp_cpu_load.items():
        status.append(f"{value}")
        perf.append(f"{key}={value}")
        if value is None:
            compare = sys.maxsize
        elif float(value)>compare:
            compare = float(value)
    nagios_status("CPU", compare, warning, critical,
      f"load: {', '.join(status)}|{'; '.join(perf)}"
    )


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 = {}
    snmp_args = {}
    for opt, arg in gopts:
        opts[opt] = arg
        if opt.strip("-") in snmp_opts:
            snmp_args[opt.strip("-")] = arg
    #print(opts)

    hostname = opts.get("-H")
    if "-J" in opts:
        json_path = opts["-J"]
        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 '-S' in opts:
        print(json.dumps(snmp_args, indent=2))
        quit()

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

    command = opts.get('-e')

    if command=="load" or command=="cpu":
        show_load(snmp.cpu_load(), warning, critical)
    elif command=="mem":
        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 command=="disk" and '-p' in opts:
        part = opts['-p']
        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 command=="disk":
        for key, value in snmp.storage().items():
            print(f"{key}: {value}")
    elif command=="net" and '-i'  in opts:
        ifname = opts['-i']
        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 command=="net":
        for key, value in snmp.network().items():
            print(f"{key}")
    elif command=="cisco_load" or command=="cisco_cpu":
        show_load(snmp.cisco_load(), warning, critical)
    else:
        print("Missing parameter: -e")
        sys.exit(5)
