#! /usr/bin/env python

#######################################################################
# Copyright (C) 2013 VMWare, Inc.
# All Rights Reserved
#######################################################################

# Watchdog daemon that starts and stops the various daemons that make up the
# netdumper service.  The watchdog will periodically check for the
# existence of each pid and do a restart if the subdaemon exited.
#
# The current set of subdaemons are:
#
#  vmware-netdumper - The netdump server.
#  vmware-netdumper-webserver - The netdump webserver.

import os
import sys
import time
import signal
import logging
import logging.handlers

LOOPING = True

POLL_INTERVAL = 5

RUN_DIR = '/var/run/vmware-netdumper'

def noop(_sig, _fr):
    return

def stopLooping(_sig, _fr):
    '''Signal handler that signals the main loop to stop.'''
    global LOOPING
    LOOPING = False

class DaemonProcess(object):
    '''Manages a daemon process in linux.'''

    # The number of seconds to wait before getting more aggressive with killing
    # processes.
    KILL_TIMEOUT = 3

    def __init__(self, name):
        self.name = name
        self.pidFileName = os.path.join(RUN_DIR, "%s.pid" % name)
        logger.info("pid file: %s, %s" % (name, self.pidFileName))
        self.pid = None

    def setPid(self, pid):
        '''Set the pid for this daemon and writes it to the pid file.'''
        self.pid = pid
        pidFile = open(self.pidFileName, 'w')
        pidFile.write("%s" % pid)
        pidFile.close()

    def recoverPid(self):
        '''Attempt to recover the pid for a daemon that was left running after
        the watchdog was killed.'''
        try:
            pidFile = open(self.pidFileName, 'r')
            pid = int(pidFile.read())
            os.kill(pid, 0)
            self.setPid(pid)
            logger.info("recovered pid from file -- %s" % self.pidFileName)
            return True
        except IOError:
            logger.debug("no existing pid file -- %s" % self.pidFileName)
        except OSError:
            logger.debug("leftover pid file -- %s" % self.pidFileName)
        except ValueError, e:
            logger.warn("invalid value in pid file -- %s, %s" % (
                    self.pidFileName, str(e)))

        return False

    def isAlive(self):
        '''Test if the daemon is alive.'''
        try:
            os.kill(self.pid, 0)
            return True
        except OSError:
            logger.warn("lost daemon %s; pid=%s" % (self.name, self.pid))
            try:
                os.remove(self.pidFileName)
            except OSError:
                logger.warn("missing pid file -- %s" % self.pidFileName)
            return False

    def killItDead(self):
        '''Try various ways to kill the daemon.'''
        if not self.pid:
            return
        for signum in [signal.SIGINT, signal.SIGTERM, signal.SIGKILL]:
            try:
                # Set an alarm to wake up from the waitpid.
                signal.alarm(self.KILL_TIMEOUT)
                os.kill(self.pid, signum)
                os.waitpid(self.pid, 0)
                try:
                    os.remove(self.pidFileName)
                except (IOError, OSError):
                    pass
                signal.alarm(0)
                break
            except OSError:
                logger.warn(
                    "daemon %s with pid %s did not respond to signal %s" % (
                        self.name, self.pid, signum))

    def __repr__(self):
        return 'DaemonProcess[name=%s; pid=%s; pidFileName=%s]' % (
            self.name, self.pid, self.pidFileName)

class ServicesWatchdog(object):
    '''Manager for the sub-daemons.  The args value should be a list of
    arguments to pass to the daemons when they are started.'''
    def __init__(self, args=None):
        self.procs = {}
        if not args:
            args = []
        self.args = args

    def _ensureDaemonIsRunning(self, daemon):
        if daemon in self.procs:
            if self.procs[daemon].isAlive():
                return
            del self.procs[daemon]

        allArgs = [daemon] + self.args
        self.procs[daemon] = DaemonProcess(os.path.basename(allArgs[0]))
        if self.procs[daemon].recoverPid():
            return

        try:
            self.procs[daemon].setPid(
                os.spawnv(os.P_NOWAIT, daemon, allArgs))
            logger.info("started %s; pid=%s" % (
                    allArgs, self.procs[daemon]))
        except OSError, e:
            logger.error("could not start daemon %s -- %s" % (
                    daemon, str(e)))

    def start(self):
        '''Ensure that all the daemons are started.'''
        for daemon in self.DAEMONS:
            self._ensureDaemonIsRunning(daemon)

    def shutdown(self):
        '''Shutdown all of the daemons.'''
        for daemon in self.DAEMONS:
            if daemon in self.procs:
                self.procs[daemon].killItDead()
                del self.procs[daemon]

class NetdumperMainService(ServicesWatchdog):
    DAEMONS = [
        "/usr/sbin/vmware-netdumper",
        ]

class NetdumperWebService(ServicesWatchdog):
    DAEMONS = [
          "/usr/lib/vmware-netdumper/webserver/vmware-netdumper-webserver",
          ]


logfile = "/var/log/vmware/netdumper/watchdog.log"
logger = logging.getLogger('netdump.watchdog')
logger.setLevel(logging.DEBUG)

rfHandler = logging.handlers.RotatingFileHandler(logfile,
                                                 maxBytes=1048576,
                                                 backupCount=10)
logformatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
rfHandler.setFormatter(logformatter)

logger.addHandler(rfHandler)

def main(args):
    activeServers = {}

    errorLogFd = os.open("/var/log/vmware/netdumper/watchdog-error.log",
                         os.O_CREAT|os.O_WRONLY|os.O_APPEND)
    os.dup2(errorLogFd, sys.__stdout__.fileno())
    os.dup2(errorLogFd, sys.__stderr__.fileno())

    signal.signal(signal.SIGALRM, noop)
    signal.signal(signal.SIGTERM, stopLooping)

    if os.fork() != 0:
        return 0

    watchDogPidFile = open(os.path.join(RUN_DIR, "watchdog.pid"), "w")
    watchDogPidFile.write("%s" % os.getpid())
    watchDogPidFile.close()

    try:
        logger.info("Netdumper watchdog starting...")

        mainsvc = NetdumperMainService(args)
        websvc = NetdumperWebService()

        while LOOPING:
            mainsvc.start()
            websvc.start()

            try:
                rc = os.waitpid(0, os.WNOHANG)
            except OSError, e:
                pass
            time.sleep(POLL_INTERVAL)
    except:
        logger.exception("Unknown Error")
    finally:
        mainsvc.shutdown()
        websvc.shutdown()
        os.remove(os.path.join(RUN_DIR, "watchdog.pid"))
        logger.info("watchdog exiting")

    return 0

if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
