#! /usr/bin/python
# Copyright 2019 Cumulus Networks, Inc.  All Rights reserved.
#
# This software is subject to the Cumulus Networks End User License Agreement
# available at the following locations:
#
# Internet: https://cumulusnetworks.com/downloads/eula/latest/view/
# Cumulus Linux systems: /usr/share/cumulus/EULA.txt
import subprocess
import ipaddr
import os
import sys
import logging
import socket
from configobj import ConfigObj


def error_handler(func):
    def wrapper(*args, **kwargs):
        try:
            print('')
            func(*args, **kwargs)
        except ValueError as ex:
            print('FAIL')
            print(ex)
            return False
        else:
            print('PASS')
            return True

    return wrapper


class OptaCheck(object):
    RC_FAIL = 1
    RC_SUCCESS = 0
    ROOT_USER = 0

    def __init__(self):
        self.proc_name = 'opta-check'
        logger = logging.getLogger(self.proc_name)
        if len(logger.handlers):
            # Logger already configured
            self.logger = logger
        else:
            handler = logging.FileHandler(filename='/var/log/opta-check.log')
            formatter = logging.Formatter(
                '%(asctime)s {} {}[%(process)d]: %(levelname)s: %(message)s'.format(socket.gethostname(),
                                                                                    self.proc_name))
            handler.setLevel(logging.INFO)
            handler.setFormatter(formatter)
            logger.setLevel(logging.INFO)
            logger.addHandler(handler)
            self.logger = logger

        if self.is_cloud_appliance():
            self.is_cloud_app = True
            msg = 'Detected NetQ Cloud Appliance'
            print(msg)
            self.logger.info(msg)
        else:
            self.is_cloud_app = False

        self.logger.info('running opta-check...')

    @staticmethod
    def is_opta():
        return os.path.isfile('/etc/app-release')

    @staticmethod
    def is_cloud_appliance():
        conf = ConfigObj('/etc/app-release')
        return(conf.get('APPLIANCE_NAME') == "NetQ Cloud Appliance")

    @staticmethod
    def is_vmware_vm():
        out = subprocess.check_output('dmidecode -s system-manufacturer'.split())
        if out.strip() == "VMware, Inc.":
            return True
        else:
            return False

    @error_handler
    def interface_check(self, ifname):
        msg = 'checking connectivity to {}...'.format(ifname)
        print(msg)
        self.logger.info(msg)

        out = ''
        try:
            out = subprocess.check_output('ip -4 -br addr show {}'.format(ifname).split())
        except subprocess.CalledProcessError:
            msg = 'undefined interface {}'.format(ifname)
            self.logger.warning(msg)
            print('WARNING: {}'.format(msg))
        finally:
            if len(out):
                ip = ipaddr.IPNetwork(out.split()[2])
                try:
                    subprocess.check_output('ping {} -c 3'.format(ip.ip).split())
                except subprocess.CalledProcessError:
                    msg = 'IP address {} is not reachable'.format(ip.ip)
                    self.logger.error(msg)
                    raise ValueError('ERROR: {}'.format(msg))
                else:
                    msg = 'IP address {} is reachable'.format(ip.ip)
                    self.logger.info(msg)
                    print('INFO: {}'.format(msg))
                    return True
            else:
                msg = 'IP address not assigned to {}'.format(ifname)
                print('WARNING: {}'.format(msg))
                self.logger.warning(msg)
                return True

    @error_handler
    def kube_ip_change(self):

        for path in ['/home/cumulus/.kube/config', '/etc/kubernetes/admin.conf']:
            msg = 'Checking for IP address change since last orchestrated in {}'.format(path)
            print(msg)
            self.logger.info(msg)
            if not os.path.isfile(path):
                msg = '{} does not exist'.format(path)
                self.logger.error(msg)
                raise ValueError('ERROR: {}'.format(msg))

            config_out = subprocess.check_output('grep server {}'.format(path).split())
            if len(config_out):
                kube_config_ip = config_out.replace('/', '').split(':')[2]
                for iface in [ 'eth0', 'eth1']:
                    try:
                        out = subprocess.check_output('ip -4 -br addr show {}'.format(iface).split())
                        ip = str(ipaddr.IPNetwork(out.split()[2]).ip)
                    except subprocess.CalledProcessError:
                        msg = 'No IP address found for interace: {}'.format(iface)
                        print('Error: {}'.format(msg))
                        self.logger.info(msg)
                        continue
                    if ip != kube_config_ip:
                        msg = '{} IP {} and api-server IP {} do not match'.format(iface, ip, kube_config_ip)
                        print('WARNING: {}'.format(msg))
                        self.logger.warning(msg)
                    else:
                        msg = '{} matches IP address of {}'.format(ip, iface)
                        print('INFO: {}'.format(msg))
                        self.logger.info(msg)
                        break
                else:
                    msg = 'ERROR: {} does not match with either eth0/eth1'.format(kube_config_ip)
                    self.logger.error(msg)
                    raise ValueError(msg)
            else:
                msg = 'server section is not defined in {}'.format(path)
                self.logger.error(msg)
                raise ValueError('ERROR: {}'.format(msg))
        return True

    @error_handler
    def default_route(self):
        msg = 'checking default route...'
        print(msg)
        self.logger.info(msg)
        out = subprocess.check_output('ip route show default 0.0.0.0/0'.split())
        if len(out):
            return True
        else:
            msg = 'no default route found'
            self.logger.error(msg)
            raise ValueError('ERROR: {}'.format(msg))

    @error_handler
    def ram_requirements(self):
        msg = 'checking RAM requirements...'
        print(msg)
        self.logger.info(msg)
        path = '/proc/meminfo'
        if not os.path.isfile(path):
            msg = '{} does not exist'.format(path)
            self.logger.error(msg)
            raise ValueError('ERROR: {}'.format(msg))
        out = subprocess.check_output('grep MemTotal: {}'.format(path).split())
        mem_total = out.split()[1]

        if self.is_cloud_app:
           min_mem = 8022214
           min_human = 8
        else:
           min_mem = 65502188
           min_human = 64

        gb_in_kb = 1048576

        if int(mem_total) < min_mem:
            msg = 'minimum of {} GB RAM required but {} GB RAM detected'.format(min_human, int(mem_total) / gb_in_kb)
            self.logger.error(msg)
            raise ValueError('ERROR: {}'.format(msg))
        else:
            msg = 'detected {} GB RAM'.format(int(mem_total) / gb_in_kb)
            self.logger.info(msg)
            print('INFO: {}'.format(msg))
            return True

    @error_handler
    def cpu_requirements(self):
        msg = 'checking CPU requirements...'
        print(msg)
        self.logger.info(msg)
        out = subprocess.check_output('nproc'.split())
        if self.is_cloud_app:
           min_cores = 4
        else:
           min_cores = 8
        num_cores = int(out)
        if num_cores < min_cores:
            msg = 'minimum of {} CPU cores required but {} detected'.format(min_cores, num_cores)
            self.logger.error(msg)
            raise ValueError('ERROR: {}'.format(msg))
        else:
            msg = 'detected {} CPU cores'.format(num_cores)
            print('INFO: {}'.format(msg))
            self.logger.info(msg)
            return True

    @error_handler
    def check_drivers(self):
        if self.is_vmware_vm():

            msg = "VMWare VM is detected. Checking for vmw_pvscsi driver."
            print(msg)
            self.logger.info(msg)

            out = subprocess.check_output('lsmod'.split())

            if 'vmw_pvscsi' in out:
                msg = "Detected vmw_pvsci driver."
                print('INFO: {}'.format(msg))
                self.logger.info(msg)
                return True
            else:
                msg = 'Did not find the vmw_pvscsi driver enabled on this NetQ VM.'
                msg += ' Please re-install the NetQ VM on ESXi server.'
                self.logger.error(msg)
                raise ValueError('ERROR: {}'.format(msg))

    #TODO: implement hardware checks later
    # @error_handler
    # def ssd_check(self):
    #     print('checking for SSD...')
    #     out = subprocess.check_output('lsblk -d -o name,rota'.split())
    #     hdd_names = []
    #     for disk in out.split('\n')[1:-1]:
    #         name, hdd = tuple(disk.split())
    #         if int(hdd) == 1:
    #             hdd_names.append(name)
    #     if len(hdd_names):
    #         print('WARNING: the following disks are HDD: {}\n'
    #               '         running on a SSD is recommended'.format(', '.join(hdd_names)))
    #     else:
    #         print('PASS')

    def run(self):
        checks = [self.interface_check('eth0'), self.interface_check('eth1'), self.kube_ip_change(),
                  self.default_route(), self.ram_requirements(), self.cpu_requirements(), 
                  self.check_drivers()]

        fail_count = 0
        all_passed = True
        for result in checks:
            if not result:
                fail_count += 1
            all_passed = all_passed and result

        print('-------')
        print('SUMMARY')
        msg = 'running on a SSD is recommended'
        print('INFO: {}'.format(msg))
        self.logger.info(msg)
        if all_passed:
            msg = 'ALL CHECKS PASSED'
            print('INFO: {}'.format(msg))
            self.logger.info('SUMMARY: {}'.format(msg))
            return OptaCheck.RC_SUCCESS
        else:
            msg = '{} CHECKS FAILED'.format(fail_count)
            print('ERROR: {}'.format(msg))
            self.logger.error('SUMMARY: {}'.format(msg))
            return OptaCheck.RC_FAIL


if __name__ == '__main__':
    if os.getuid() != OptaCheck.ROOT_USER:
        msg = 'please run this command with sudo'
        print('ERROR: {}'.format(msg))
        sys.exit(OptaCheck.RC_FAIL)

    opta_check = OptaCheck()
    if not OptaCheck.is_opta():
        msg = 'please run on opta'
        opta_check.logger.error(msg)
        print('ERROR: {}'.format(msg))
        sys.exit(OptaCheck.RC_FAIL)

    sys.exit(opta_check.run())
