#! /usr/bin/python
# Copyright 2017 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


from __future__ import absolute_import

import atexit
import argparse
import bisect
import functools
import itertools
import json
import os
import socket
import subprocess
import sys
import traceback
from datetime import datetime
from operator import itemgetter
from threading import Thread, Event
from daemon import DaemonContext
import signal
import re
from cStringIO import StringIO
import pwd
import struct

import tabulate
from network_docopt import (
    NetworkDocopt,
    get_network_docopt_info
)

from netq_lib.common import utils
from netq_lib.common.enums import (
    BgColors,
    DiffState,
    Heartbeat,
    ConfKw,
    ServiceStatus,
    ServiceStatusString
)
from netq_apps.lib.netq import (
    NetQ,
)
from netq_lib.common.service import Service
from netq_apps.lib.netq import NetQ, NetQException
from netq_lib.common.enums import ConfKw
from netq_apps.cmd.netq import NetqCLI


_JSON = None
_EMPTY_JSON = '[{}]'
_RC_SUCCESS = 0
_RC_FAIL = 1
_PROC_NAME = 'netqd'

class Capturing(list):

    def __enter__(self):
        self._stdout = sys.stdout
        sys.stdout = self._stringio = StringIO()
        return self

    def __exit__(self, *args):
        self.extend(self._stringio.getvalue().splitlines())
        sys.stdout = self._stdout


class NetqDaemon(object):
    def __init__(self):
        self.shutdown_event = Event()
        self.users_with_edit = {}
        self.groups_with_edit = {}
        self.users_with_show = {}
        self.groups_with_show = {}
        self.cli = None
        self._logger = None

    def signal_handler(self, signal, frame):
        # signal handlers are initialized before _logger is initialized in main
        # A SIGNAL could sneak through in that window.
        if self._logger:
            self._logger.info("received SIGINT or SIGTERM")
        self.shutdown_event.set()

    def reload_handler(self, signal, frame):
        self._logger.info("received SIGHUP")
        self.init_cli()

    def tx_reply(self, connection, reply):

        if reply is None:
            reply = ''

        if not reply.endswith('\n'):
            reply += '\n'

        # TX the reply to the 'net' instance
        try:
            connection.send(reply)
        except socket.error:
            self._logger.info("TX reply failed")

        connection.close()
        # log.info("TXed reply")

    def init_cli(self):
        self.cli = NetqCLI(self.daemon, self._logger)

    def get_pid_uid_gid(self, uds):
        """
        Obtain the effective user and group IDs of the process on the other end
        of a socket. SO_PEERCRED is used so the information returned is
        trustworthy. (We are OK with root saying they are somebody less
        privileged.)
        """
        pid = None
        uid = None
        gid = None

        try:
            # rely on pid, uid, gid being None if family is not AF_UNIX
            if uds.family != socket.AF_UNIX:
                return (pid, uid, gid)
        except AttributeError:
            pass

        try:
            credentials = uds.getsockopt(socket.SOL_SOCKET, self.SO_PEERCRED,
                                         struct.calcsize('3i'))
            pid, uid, gid = struct.unpack('3i', credentials)

            if pid == -1:
                pid = None

            if uid == -1:
                uid = None

            if gid == -1:
                gid = None

        except Exception as e:
            self._logger.error("socket get credentials failed: {0}".format(e))

        return (pid, uid, gid)

    def __delpid(self):
        """ Removes the pidfile when the process exits gracefully.
        """
        try:
            os.remove('/var/run/%s.pid' % _PROC_NAME)
        except OSError:
            self._logger.exception('Unable to remove pid file on exit')

    def main(self, daemon, conf):
        self._logger = utils.get_logger(_PROC_NAME, conf)
        self.__pidfd = utils.Pidfile('/var/run/%s.pid' % _PROC_NAME,
                                     _PROC_NAME)
        uds = None
        try:
            # Log the process ID upon start-up.
            try:
                netqd_pid = os.getpid()
                # Make sure that I am not already running.
                self.__pidfd.lock()
                self.__pidfd.write(netqd_pid)
                atexit.register(self.__delpid)
                self._logger.info('Starting (pid %d) ...', netqd_pid)
                atexit.register(self._logger.info,
                                'Terminating (pid %d)' % netqd_pid)
            except OSError as e:
                self._logger.exception(e)
                netqd_pid = 'Error'
                caught_exception = True

            uds = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
            self._logger.info('netqd started with process ID (pid) {0}.'.format(netqd_pid))

            try:
                self.SO_PEERCRED = socket.SO_PEERCRED
                self.SO_PASSCRED = socket.SO_PASSCRED
            except AttributeError:
                # powerpc is the only non-generic we care about. alpha, mips,
                # sparc, and parisc also have non-generic values.
                machine = os.uname()[4]
                if re.search(r"""^(ppc|powerpc)""", machine):
                    self.SO_PASSCRED = 20
                    self.SO_PEERCRED = 21
                else:
                    self.SO_PASSCRED = 16
                    self.SO_PEERCRED = 17

            self.server_address = '/var/run/netqd/uds'
            if os.path.exists(self.server_address):
                os.remove(self.server_address)

            self.daemon = daemon
            try:
                self.init_cli()
            except Exception as ex:
                self._logger.info('An exception of type %s occured' % type(ex))
                pass

            uds.bind(self.server_address)
            uds.setsockopt(socket.SOL_SOCKET, self.SO_PASSCRED, 1)
            uds.listen(1)
            os.chmod(self.server_address, 0777)

            while True:
                if self.shutdown_event.is_set():
                    self._logger.debug('Shutdown signal RXed.  '
                                       'Breaking out of the loop.')
                    break

                try:
                    (connection, _) = uds.accept()
                except socket.error as e:
                    if isinstance(e.args, tuple) and e[0] == 4:
                        # 4 is 'Interrupted system call', a.k.a. SIGINT.
                        # The user wants to stop netqd.
                        self._logger.info('socket.accept() caught signal, '
                                          'starting shutdown')
                        self.shutdown_event.set()
                        continue  # Avoid further reference to "connection"
                    else:
                        self._logger.info("netqd socket %s hit an error\n%s" %
                                          (self.server_address, e))
                        uds.close()
                        uds = None

                    if os.path.exists(self.server_address):
                        os.remove(self.server_address)

                    uds = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
                    uds.bind(self.server_address)
                    uds.setsockopt(socket.SOL_SOCKET, self.SO_PASSCRED, 1)
                    uds.listen(1)
                    continue

                # RX the request from the 'net' instance.
                data = connection.recv(4096)
                try:
                    argv = json.loads(data)  # argv is a string
                except ValueError as ex:
                    self._logger.exception(ex)
                    self.cli.re_init()
                    continue

                len_argv = len(argv)
                (pid, uid, gid) = self.get_pid_uid_gid(connection)
                user = pwd.getpwuid(uid)[0]

                reply = None

                # The logs would be pretty chatty if we logged all of the TAB
                # complete calls, so make those log.debug() and
                # everything else info().
                if '--completions' in argv:
                    self._logger.debug("RXed: command '%s'" % (' '.join(argv)))
                else:
                    self._logger.info("RXed: command '%s'" % (' '.join(argv)))

                # bash model
                if '--completions' in argv:
                    # the last 2 args are related to completions
                    sys.argv = argv[:-2]
                    (print_options, ended_with_space, argv) = (
                        get_network_docopt_info(argv))
                    options_style = 'bash'

                # fish model
                elif '--fish-completion' in argv:
                    # For "foo sh <TAB>" argv will be
                    # [foo', '--fish-completion', 'foo show ']
                    # the last 2 args are related to completions
                    sys.argv = argv[:-2]
                    last_arg = argv[-1]

                    if last_arg.endswith(' '):
                        ended_with_space = True
                    else:
                        ended_with_space = False

                    argv = last_arg.strip().split()
                    print_options = True
                    options_style = 'fish'
                else:
                    print_options = False

                reply = []
                rc = 0
                with Capturing() as reply:
                    if print_options:
                        self.cli.print_options(ended_with_space, options_style,
                                               argv)
                    else:
                        # This is where the magic happens
                        # (the command gets executed).
                        new_argv = []
                        for arg in argv:
                            new_argv.append(arg.encode('ascii', 'ignore'))
                        try:
                            rc = self.cli.run(new_argv, uid, gid)
                        except (Exception, NetQException, RuntimeError) as ex:
                            self._logger.exception(ex)
                            if argv[-1] == 'json':
                                print json.dumps([{'Error': '{}'.format(ex)}],
                                                 indent=4)
                            else:
                                print utils.color_wrap(BgColors.RED, str(ex))
                            rc = 1
                output = {'rc': rc, 'output': '\n'.join(reply)}
                self.tx_reply(connection, json.dumps(output))

                self.cli.re_init()

            self._logger.info(
                'netqd is stopping with process ID (pid) {0}.'.format(netqd_pid))

        except Exception as e:
            self._logger.exception(e)
            caught_exception = True

            if uds:
                uds.close()
                uds = None

            if caught_exception:
                sys.exit(1)
            else:
                sys.exit(0)


if __name__ == '__main__':
    parser = argparse.ArgumentParser(prog=_PROC_NAME,
                                     description='Netq CLI Daemon')
    parser.add_argument('-s', '--server',
                        help='Add Backend server name or IP to the config file')
    parser.add_argument('-d', '--daemon', help='run as a daemon', action='store_true')

    args = parser.parse_args()
    conf_dict = utils.get_yaml_config(_PROC_NAME)
    if args.server:
        if ConfKw.BACKEND not in conf_dict:
            conf_dict[ConfKw.BACKEND] = {
                ConfKw.SERVER: args.server}
        else:
            conf_dict[ConfKw.BACKEND][ConfKw.SERVER] = (
                args.server)
        exit(utils.update_config(conf_dict))

    netqd = NetqDaemon()

    if not os.path.exists('/var/run/netqd/'):
        os.makedirs('/var/run/netqd/', mode=0755)

    if args.daemon:
        context = DaemonContext(
            working_directory='/var/run/netqd/',
            signal_map={
                signal.SIGTERM: netqd.signal_handler,
                signal.SIGINT: netqd.signal_handler,
                signal.SIGHUP: netqd.reload_handler,
            }
        )

        context.open()
        with context:
            netqd.main(True, conf_dict)

    else:
        signal.signal(signal.SIGINT, netqd.signal_handler)
        signal.signal(signal.SIGTERM, netqd.signal_handler)
        signal.signal(signal.SIGHUP, netqd.reload_handler)
        netqd.main(False, conf_dict)
