#!/usr/libexec/platform-python
#
# Copyright 2023 Tencent.com, Inc. and its affiliates. All Rights Reserved.
#
# Licensed under the MIT License. See the LICENSE accompanying this file
# for the specific language governing permissions and limitations under
# the License.
#

import base64
import errno
import hashlib
import hmac
import json
import logging
import logging.handlers
import os
import platform
import pwd
import re
import shutil
import socket
import subprocess
import sys
import time
from collections import namedtuple
from contextlib import contextmanager
from datetime import datetime, timedelta
from logging.handlers import RotatingFileHandler
from signal import SIGHUP, SIGKILL, SIGTERM

try:
    from configparser import ConfigParser, NoOptionError, NoSectionError
except ImportError:
    import ConfigParser
    from ConfigParser import NoOptionError, NoSectionError


VERSION = "1.0.0"
SERVICE = "cloudfilesystem"

CONFIG_FILE = "/etc/cfs/cfs-utils.conf"
CONFIG_SECTION = "mount-watchdog"


DEFAULT_UNKNOWN_VALUE = "unknown"
# 50ms
DEFAULT_TIMEOUT = 0.05

LOG_DIR = "/var/log/cfs"
LOG_FILE = "mount-watchdog.log"

STATE_FILE_DIR = "/var/run/cfs"
STUNNEL_PID_FILE = "stunnel.pid"

DEFAULT_NFS_PORT = "2049"
DEFAULT_STUNNEL_HEALTH_CHECK_INTERVAL_MIN = 5
DEFAULT_STUNNEL_HEALTH_CHECK_TIMEOUT_SEC = 30


Mount = namedtuple(
    "Mount", ["server", "mountpoint", "type", "options", "freq", "passno"]
)

# Unmount difference time in seconds
UNMOUNT_DIFF_TIME = 30

# Default unmount count for consistency
DEFAULT_UNMOUNT_COUNT_FOR_CONSISTENCY = 5

SYSTEM_RELEASE_PATH = "/etc/system-release"
OS_RELEASE_PATH = "/etc/os-release"
STUNNEL_INSTALLATION_MESSAGE = "Please install stunnel"


def fatal_error(user_message, log_message=None):
    if log_message is None:
        log_message = user_message

    sys.stderr.write("%s\n" % user_message)
    logging.error(log_message)
    sys.exit(1)


def get_boolean_config_item_value(
    config, config_section, config_item, default_value, emit_warning_message=False
):
    warning_message = None
    if not config.has_section(config_section):
        warning_message = (
            "Warning: config file does not have section %s." % config_section
        )
    elif not config.has_option(config_section, config_item):
        warning_message = (
            "Warning: config file does not have %s item in section %s."
            % (config_item, config_section)
        )

    if warning_message:
        if emit_warning_message:
            sys.stdout.write(
                "%s. You should be able to find a new config file in the same folder as current config file %s. "
                "Consider update the new config file to latest config file. Use the default value [%s = %s]."
                % (warning_message, CONFIG_FILE, config_item, default_value)
            )
        return default_value
    return config.getboolean(config_section, config_item)


def bootstrap_logging(config, log_dir=LOG_DIR):
    raw_level = config.get(CONFIG_SECTION, "logging_level")
    levels = {
        "debug": logging.DEBUG,
        "info": logging.INFO,
        "warning": logging.WARNING,
        "error": logging.ERROR,
        "critical": logging.CRITICAL,
    }
    level = levels.get(raw_level.lower())
    level_error = False

    if not level:
        # delay logging error about malformed log level until after logging is configured
        level_error = True
        level = logging.INFO

    max_bytes = config.getint(CONFIG_SECTION, "logging_max_bytes")
    file_count = config.getint(CONFIG_SECTION, "logging_file_count")

    handler = RotatingFileHandler(
        os.path.join(log_dir, LOG_FILE), maxBytes=max_bytes, backupCount=file_count
    )
    handler.setFormatter(
        logging.Formatter(
            fmt="%(asctime)s - %(levelname)s - %(message)s",
            datefmt="%Y-%m-%d %H:%M:%S %Z",
        )
    )

    logger = logging.getLogger()
    logger.setLevel(level)
    logger.addHandler(handler)

    if level_error:
        logging.error(
            'Malformed logging level "%s", setting logging level to %s',
            raw_level,
            level,
        )


def parse_options(options):
    opts = {}
    for o in options.split(","):
        if "=" in o:
            k, v = o.split("=")
            opts[k] = v
        else:
            opts[o] = None
    return opts


def get_file_safe_mountpoint(mount):
    mountpoint = os.path.abspath(mount.mountpoint).replace(os.sep, ".")
    if mountpoint.startswith("."):
        mountpoint = mountpoint[1:]

    opts = parse_options(mount.options)
    if "port" not in opts:
            # /proc/mounts provides a list of all mounts in use by the system (including the mount options used).
            # In the case of tls mount: stunnel establishes a localhost port connection in order to listen on the requests,
            # and then send packets further to the server:2049. If the port is 2049 which is the default nfs port,
            # /proc/mounts will not display the port number in the options information, thus watchdog process will not treat
            # the mount as CFS mount and won't restart the killed stunnel which cause the mount hang.
            # So, tlsport=2049 is being added here by appending with the mountpoint.
            # Putting a default port 2049 to fix the Stunnel process being killed issue.
            opts["port"] = DEFAULT_NFS_PORT

    return mountpoint + "." + opts["port"]


def get_current_local_nfs_mounts(mount_file="/proc/mounts"):
    """
    Return a dict of the current NFS mounts for servers running on localhost, keyed by the mountpoint and port as it
    appears in CFS watchdog state files.
    """
    mounts = []

    with open(mount_file) as f:
        for mount in f:
            try:
                mounts.append(Mount._make(mount.strip().split()))
            except Exception as e:
                # Make sure nfs mounts being skipped are made apparent
                if " nfs4 " in mount:
                    logging.warning(
                        'Watchdog ignoring malformed nfs4 mount "%s": %s', mount, e
                    )
                else:
                    logging.debug(
                        'Watchdog ignoring malformed mount "%s": %s', mount, e
                    )

    mounts = [m for m in mounts if m.server.startswith("127.0.0.1") and "nfs" in m.type]

    mount_dict = {}
    for m in mounts:
        safe_mnt = get_file_safe_mountpoint(m)
        if safe_mnt:
            mount_dict[safe_mnt] = m

    return mount_dict


def get_state_files(state_file_dir):
    """
    Return a dict of the absolute path of state files in state_file_dir,
    keyed by the mountpoint and port portion of the filename.
    """
    state_files = {}

    if os.path.isdir(state_file_dir):
        for sf in os.listdir(state_file_dir):
            if not sf.startswith("cfs-") or os.path.isdir(
                os.path.join(state_file_dir, sf)
            ):
                continue

            # This translates the state file name "cfs-10.0.0.1.mnt.nfs.12345"
            # into file-safe mountpoint "mnt.nfs.12345"
            first_period = sf.find(".")
            first_period = sf.find(".", first_period + 1)
            first_period = sf.find(".", first_period + 1)
            first_period = sf.find(".", first_period + 1)
            mount_point_and_port = sf[first_period + 1 :]
            logging.debug(
                'Translating "%s" into mount point and port "%s"',
                sf,
                mount_point_and_port,
            )
            state_files[mount_point_and_port] = sf

    return state_files


def get_pid_in_state_dir(state_file, state_file_dir):
    """
    :param state_file: The state file path, e.g. cfs-10.0.0.1.mnt.20560.
    :param state_file_dir: The state file dir path, e.g. /var/run/cfs.
    """
    state_dir_pid_path = os.path.join(
        state_file_dir, state_file + ".dir", STUNNEL_PID_FILE
    )
    if os.path.exists(state_dir_pid_path):
        with open(state_dir_pid_path) as f:
            return f.read()
    return None


def is_mount_stunnel_proc_running(state_pid, state_file, state_file_dir):
    """
    Check whether a given stunnel process id in state file is running for the mount. To avoid we incorrectly checking
    processes running by other applications and send signal further, the stunnel process in state file is counted as
    running iff:
    1. The pid in state file is not None.
    2. The process running with the pid is a stunnel process. This is validated through process command name.
    3. The process can be reached via os.kill(pid, 0).
    4. Every launched stunnel process will write its process id to the pid file in the mount state_file_dir, and only
       when the stunnel is terminated this pid file can be removed. Check whether the stunnel pid file exists and its
       value is equal to the pid documented in state file. This step is to make sure we don't send signal later to any
       stunnel process that is not owned by the mount.

    :param state_pid: The pid in state file.
    :param state_file: The state file path, e.g. cfs-10.0.0.1.mnt.20560.
    :param state_file_dir: The state file dir path, e.g. /var/run/cfs.
    """
    if not state_pid:
        logging.debug("State pid is None for %s", state_file)
        return False

    process_name = check_process_name(state_pid)
    if not process_name or "stunnel" not in str(process_name):
        logging.debug(
            "Process running on %s is not a stunnel process, full command: %s.",
            state_pid,
            str(process_name) if process_name else "",
        )
        return False

    if not is_pid_running(state_pid):
        logging.debug(
            "Stunnel process with pid %s is not running anymore for %s.",
            state_pid,
            state_file,
        )
        return False

    pid_in_stunnel_pid_file = get_pid_in_state_dir(state_file, state_file_dir)
    # To avoid the healthy stunnel established by those version to be treated as not running due to the missing pid file, which can result in stunnel being constantly restarted,
    # assuming the stunnel is still running even if the stunnel pid file does not exist.
    if not pid_in_stunnel_pid_file:
        logging.debug(
            "Pid file of stunnel does not exist for %s. It is possible that the stunnel is no longer running or the mount was mounted using an older version cfs-utils (<1.32.2). Assuming the stunnel with pid %s is still running.",
            state_file,
            state_pid,
        )

    elif int(state_pid) != int(pid_in_stunnel_pid_file):
        logging.warning(
            "Stunnel pid mismatch in state file (pid = %s) and stunnel pid file (pid = %s). Assuming the "
            "stunnel is not running.",
            int(state_pid),
            int(pid_in_stunnel_pid_file),
        )
        return False

    logging.debug("TLS tunnel for %s is running with pid %s", state_file, state_pid)
    return True


def is_pid_running(pid):
    if not pid:
        return False
    try:
        os.kill(pid, 0)
        return True
    except OSError:
        return False


def get_system_release_version():
    try:
        with open(SYSTEM_RELEASE_PATH) as f:
            return f.read().strip()
    except IOError:
        logging.debug("Unable to read %s", SYSTEM_RELEASE_PATH)

    try:
        with open(OS_RELEASE_PATH) as f:
            for line in f:
                if "PRETTY_NAME" in line:
                    return line.split("=")[1].strip()
    except IOError:
        logging.debug("Unable to read %s", OS_RELEASE_PATH)

    return DEFAULT_UNKNOWN_VALUE


def find_command_path(command, install_method):
    env_path = (
        "/sbin:/usr/sbin:/usr/local/sbin:/root/bin:/usr/local/bin:/usr/bin:/bin"
    )
    os.putenv("PATH", env_path)

    try:
        path = subprocess.check_output(["which", command])
        return path.strip().decode()
    except subprocess.CalledProcessError as e:
        fatal_error(
            "Failed to locate %s in %s - %s" % (command, env_path, install_method), e
        )


def start_tls_tunnel(child_procs, state, state_file_dir, state_file):
    # launch the tunnel in a process group so if it has any child processes, they can be killed easily
    command = state["cmd"]
    logging.info('Starting TLS tunnel: "%s"', " ".join(command))

    tunnel = None
    try:
        tunnel = subprocess.Popen(
            command,
            stdout=subprocess.DEVNULL,
            stderr=subprocess.DEVNULL,
            preexec_fn=os.setsid,
            close_fds=True,
        )
    except FileNotFoundError as e:
        logging.warning("Watchdog failed to start stunnel due to %s", e)        

    if tunnel is None or not is_pid_running(tunnel.pid):
        fatal_error(
            "Failed to initialize TLS tunnel for %s" % state_file,
            "Failed to start TLS tunnel.",
        )

    logging.info("Started TLS tunnel, pid: %d", tunnel.pid)

    child_procs.append(tunnel)
    return tunnel.pid


def clean_up_mount_state(state_file_dir, state_file, pid, mount_state_dir=None):
    send_signal_to_running_stunnel_process_group(
        pid, state_file, state_file_dir, SIGTERM
    )
    cleanup_mount_state_if_stunnel_not_running(
        pid, state_file, state_file_dir, mount_state_dir
    )


def cleanup_mount_state_if_stunnel_not_running(
    pid, state_file, state_file_dir, mount_state_dir
):
    if is_mount_stunnel_proc_running(pid, state_file, state_file_dir):
        logging.info("TLS tunnel: %d is still running, will retry termination", pid)
    else:
        if not pid:
            logging.info("TLS tunnel has been killed, cleaning up state")
        else:
            logging.info("TLS tunnel: %d is no longer running, cleaning up state", pid)
        state_file_path = os.path.join(state_file_dir, state_file)
        with open(state_file_path) as f:
            state = json.load(f)

        for f in state.get("files", list()):
            logging.debug("Deleting %s", f)
            try:
                os.remove(f)
                logging.debug("Deleted %s", f)
            except OSError as e:
                if e.errno != errno.ENOENT:
                    raise

        os.remove(state_file_path)

        if mount_state_dir is not None:
            mount_state_dir_abs_path = os.path.join(state_file_dir, mount_state_dir)
            if os.path.isdir(mount_state_dir_abs_path):
                shutil.rmtree(mount_state_dir_abs_path)
            else:
                logging.debug(
                    "Attempt to remove mount state directory %s failed. Directory is not present.",
                    mount_state_dir_abs_path,
                )


def rewrite_state_file(state, state_file_dir, state_file):
    tmp_state_file = os.path.join(state_file_dir, "~%s" % state_file)
    logging.debug(
        "Rewriting state file: writing "
        + str(len(json.dumps(state)))
        + " characters into the state file "
        + str(tmp_state_file)
    )
    with open(tmp_state_file, "w") as f:
        json.dump(state, f)

    os.rename(tmp_state_file, os.path.join(state_file_dir, state_file))


def mark_as_unmounted(state, state_file_dir, state_file, current_time):
    logging.debug("Marking %s as unmounted at %d", state_file, current_time)
    state["unmount_time"] = current_time

    rewrite_state_file(state, state_file_dir, state_file)

    return state


def restart_tls_tunnel(child_procs, state, state_file_dir, state_file):
    new_tunnel_pid = start_tls_tunnel(child_procs, state, state_file_dir, state_file)
    state["pid"] = new_tunnel_pid

    logging.debug("Rewriting %s with new pid: %d", state_file, new_tunnel_pid)
    rewrite_state_file(state, state_file_dir, state_file)


def check_cfs_mounts(
    config,
    child_procs,
    unmount_grace_period_sec,
    unmount_count_for_consistency,
    state_file_dir=STATE_FILE_DIR,
):
    nfs_mounts = get_current_local_nfs_mounts()
    logging.debug("Current local NFS mounts: %s", list(nfs_mounts.values()))

    state_files = get_state_files(state_file_dir)
    logging.debug(
        'Current state files in "%s": %s', state_file_dir, list(state_files.values())
    )

    for mount, state_file in state_files.items():
        state_file_path = os.path.join(state_file_dir, state_file)
        with open(state_file_path) as f:
            try:
                state = json.load(f)
            except ValueError:
                logging.exception("Unable to parse json in %s", state_file_path)
                continue

        current_time = time.time()
        if "unmount_time" in state:
            if state["unmount_time"] + unmount_grace_period_sec < current_time:
                logging.info("Unmount grace period expired for %s", state_file)
                clean_up_mount_state(
                    state_file_dir,
                    state_file,
                    state.get("pid"),
                    state.get("mountStateDir"),
                )
        elif mount not in nfs_mounts and (
            mount[: mount.rindex(".")] not in nfs_mounts
        ):
            # Wait 30 seconds before deciding mount no longer exists to prevent race condition
            # of watchdog's reads of nfs mounts and state files.
            if current_time - state.get("mount_time", 0) > UNMOUNT_DIFF_TIME:
                # Ensure we have consistent unmount reads for at least 5 times by default.
                if state.get("unmount_count", 0) > unmount_count_for_consistency:
                    logging.info('No mount found for "%s"', state_file)
                    state = mark_as_unmounted(
                        state, state_file_dir, state_file, current_time
                    )
                else:
                    state["unmount_count"] = state.get("unmount_count", 0) + 1
                    rewrite_state_file(state, state_file_dir, state_file)

        else:
            # Set unmount count to 0 if there were inconsistent reads
            state["unmount_count"] = 0
            rewrite_state_file(state, state_file_dir, state_file)

            if is_mount_stunnel_proc_running(
                state.get("pid"), state_file, state_file_dir
            ):
                # We have seen CFS hanging issue caused by stuck stunnel (version: 4.56) process. Apart from checking 
                # whether stunnel is running or not, we need to check whether the stunnel connection established is 
                # healthy periodically.
                #
                # The way to check the stunnel health is by `df` the mountpoint, i.e. check the file system information,
                # which will trigger a remote GETATTR on the root of the file system. Normally the command will finish
                # in 10 milliseconds, thus if the command hang for certain period (defined as 30 sec as of now), the
                # stunnel connection is likely to be unhealthy. Watchdog will kill the old stunnel process and restart
                # a new one for the unhealthy mount. The health check will run every 5 min since mount.
                #
                # Both the command hang timeout and health check interval are configurable in cfs-utils config file.
                #
                check_stunnel_health(
                    config, state, state_file_dir, state_file, child_procs, nfs_mounts
                )
            else:
                logging.warning("TLS tunnel for %s is not running", state_file)
                restart_tls_tunnel(child_procs, state, state_file_dir, state_file)


def check_stunnel_health(
    config, state, state_file_dir, state_file, child_procs, nfs_mounts
):
    if not get_boolean_config_item_value(
        config, CONFIG_SECTION, "stunnel_health_check_enabled", default_value=True
    ):
        return

    check_interval_min = get_int_value_from_config_file(
        config,
        "stunnel_health_check_interval_min",
        DEFAULT_STUNNEL_HEALTH_CHECK_INTERVAL_MIN,
    )

    current_time = time.time()

    if "mount_time" not in state:
        state["mount_time"] = current_time
        rewrite_state_file(state, state_file_dir, state_file)
        return

    # Only start to perform the stunnel health check after the check interval passed.
    if current_time - state["mount_time"] < check_interval_min * 60:
        return

    last_stunnel_check_time = (
        state["last_stunnel_check_time"] if "last_stunnel_check_time" in state else 0
    )
    if (
        last_stunnel_check_time != 0
        and current_time - last_stunnel_check_time < check_interval_min * 60
    ):
        return

    # We add this mountpoint info in the state file along with this change. It is possible for existing mounts, there
    # are no mountpoint in state file, which will cause watchdog to crash. To handle that case, we need to extract the
    # mountpoint from the state file name, and write that information to state file.
    #
    if "mountpoint" in state:
        mountpoint = state["mountpoint"]
    else:
        mountpoint = get_mountpoint_from_nfs_mounts(state_file, nfs_mounts)
        state["mountpoint"] = mountpoint
        rewrite_state_file(state, state_file_dir, state_file)

    stunnel_pid = state["pid"]
    process = subprocess.Popen(
        ["df", mountpoint],
        stdout=subprocess.DEVNULL,
        stderr=subprocess.DEVNULL,
        close_fds=True,
    )

    command_timeout_sec = get_int_value_from_config_file(
        config,
        "stunnel_health_check_command_timeout_sec",
        DEFAULT_STUNNEL_HEALTH_CHECK_TIMEOUT_SEC,
    )
    try:
        state["last_stunnel_check_time"] = current_time
        process.communicate(timeout=command_timeout_sec)
        logging.debug(
            "Stunnel [PID: %d] running for tls mount on %s passed health check.",
            stunnel_pid,
            mountpoint,
        )
        rewrite_state_file(state, state_file_dir, state_file)
    except subprocess.TimeoutExpired:
        if send_signal_to_running_stunnel_process_group(
            stunnel_pid, state_file, state_file_dir, SIGKILL
        ):
            logging.warning(
                "Connection timeout for %s after %d sec, SIGKILL has been sent to the potential unhealthy stunnel %s, "
                "restarting a new stunnel process.",
                mountpoint,
                command_timeout_sec,
                stunnel_pid,
            )
            restart_tls_tunnel(child_procs, state, state_file_dir, state_file)
        else:
            logging.warning(
                "Stunnel health check timed out for %s, stunnel [PID: %d] is not running anymore.",
                mountpoint,
                stunnel_pid,
            )
        # The child process is not killed if the timeout expires, so in order to cleanup properly, kill the child
        # process after the timeout.
        #
        process.kill()


# Retrieve the nfs mountpoint with the port information in the mount option
def get_mountpoint_from_nfs_mounts(state_file, nfs_mounts):
    search_pattern = "port={port}".format(
        port=os.path.basename(state_file).split(".")[-1]
    )
    for mount in nfs_mounts.values():
        if search_pattern in mount[3]:
            return mount[1]


def get_int_value_from_config_file(config, config_name, default_config_value):
    val = default_config_value
    try:
        value_from_config = config.get(CONFIG_SECTION, config_name)
        try:
            if int(value_from_config) > 0:
                val = int(value_from_config)
            else:
                logging.debug(
                    '%s value in config file "%s" is lower than 1. Defaulting to %d.',
                    config_name,
                    CONFIG_FILE,
                    default_config_value,
                )
        except ValueError:
            logging.debug(
                'Bad %s, "%s", in config file "%s". Defaulting to %d.',
                config_name,
                value_from_config,
                CONFIG_FILE,
                default_config_value,
            )
    except NoOptionError:
        logging.debug(
            'No %s value in config file "%s". Defaulting to %d.',
            config_name,
            CONFIG_FILE,
            default_config_value,
        )

    return val


def check_child_procs(child_procs):
    for proc in child_procs:
        proc.poll()
        if proc.returncode is not None:
            logging.warning(
                "Child TLS tunnel process %d has exited, returncode=%d",
                proc.pid,
                proc.returncode,
            )
            child_procs.remove(proc)


def parse_arguments(args=None):
    if args is None:
        args = sys.argv

    if "-h" in args[1:] or "--help" in args[1:]:
        sys.stdout.write("Usage: %s [--version] [-h|--help]\n" % args[0])
        sys.exit(0)

    if "--version" in args[1:]:
        sys.stdout.write("%s Version: %s\n" % (args[0], VERSION))
        sys.exit(0)


def assert_root():
    if os.geteuid() != 0:
        sys.stderr.write("only root can run cfs-mount-watchdog\n")
        sys.exit(1)


def read_config(config_file=CONFIG_FILE):
    try:
        p = ConfigParser.SafeConfigParser()
    except AttributeError:
        p = ConfigParser()
    p.read(config_file)
    return p


def send_signal_to_running_stunnel_process_group(
    stunnel_pid, state_file, state_file_dir, signal
):
    """
    Send a signal to the given stunnel_pid if the process running with the pid is the mount stunnel process.

    :param stunnel_pid: The pid in state file.
    :param state_file: The state file path, e.g. cfs-10.0.0.1.mnt.20560.
    :param state_file_dir: The state file dir path, e.g. /var/run/cfs.
    :param signal: OS signal send to stunnel process group, e.g. SIGHUP, SIGKILL, SIGTERM.
    """
    if is_mount_stunnel_proc_running(stunnel_pid, state_file, state_file_dir):
        process_group = os.getpgid(stunnel_pid)
        try:
            logging.info(
                "Sending signal %s(%d) to stunnel. PID: %d, group ID: %s",
                signal.name,
                signal.value,
                stunnel_pid,
                process_group,
            )
        except AttributeError:
            # In python3.4, the signal is a int object, so it does not have name and value property
            logging.info(
                "Sending signal(%s) to stunnel. PID: %d, group ID: %s",
                signal,
                stunnel_pid,
                process_group,
            )
        os.killpg(process_group, signal)
        return True
    else:
        logging.warning("TLS tunnel is not running for %s", state_file)
        return False


def subprocess_call(cmd, error_message):
    """Helper method to run shell openssl command and to handle response error messages"""
    process = subprocess.Popen(
        cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, close_fds=True
    )
    (output, err) = process.communicate()
    rc = process.poll()
    if rc != 0:
        logging.debug(
            '%s. Command %s failed, rc=%s, stdout="%s", stderr="%s"',
            error_message,
            cmd,
            rc,
            output,
            err,
        )
    else:
        return output, err


def get_utc_now():
    """
    Wrapped for patching purposes in unit tests
    """
    return datetime.utcnow()


def check_process_name(pid):
    cmd = ["cat", "/proc/{pid}/cmdline".format(pid=pid)]

    p = subprocess.Popen(
        cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, close_fds=True
    )
    return p.communicate()[0]


def clean_up_previous_stunnel_pids(state_file_dir=STATE_FILE_DIR):
    """
    Cleans up stunnel pids created by mount watchdog spawned by a previous cfs-csi-driver pod after driver restart, upgrade
    or crash. This method attempts to clean PIDs from persisted state files after cfs-csi-driver restart to
    ensure watchdog creates a new stunnel.
    """
    state_files = get_state_files(state_file_dir)
    logging.debug(
        'Persisted state files in "%s": %s', state_file_dir, list(state_files.values())
    )

    for state_file in state_files.values():
        state_file_path = os.path.join(state_file_dir, state_file)
        with open(state_file_path) as f:
            try:
                state = json.load(f)
            except ValueError:
                logging.exception("Unable to parse json in %s", state_file_path)
                continue

            try:
                pid = state["pid"]
            except KeyError:
                logging.debug("No PID found in state file %s", state_file)
                continue

            out = check_process_name(pid)

            if out and "stunnel" in str(out):
                logging.debug(
                    "PID %s in state file %s is active. Skipping clean up",
                    pid,
                    state_file,
                )
                continue

            state.pop("pid")
            logging.debug("Cleaning up pid %s in state file %s", pid, state_file)

            rewrite_state_file(state, state_file_dir, state_file)


def main():
    parse_arguments()
    assert_root()

    config = read_config()
    bootstrap_logging(config)

    child_procs = []

    if get_boolean_config_item_value(
        config, CONFIG_SECTION, "enabled", default_value=True, emit_warning_message=True
    ):
        logging.info(
            "cfs-mount-watchdog, version %s, is enabled and started", VERSION
        )
        poll_interval_sec = config.getint(CONFIG_SECTION, "poll_interval_sec")

        if config.has_option(CONFIG_SECTION, "unmount_count_for_consistency"):
            unmount_count_for_consistency = config.getint(
                CONFIG_SECTION, "unmount_count_for_consistency"
            )
        else:
            unmount_count_for_consistency = DEFAULT_UNMOUNT_COUNT_FOR_CONSISTENCY

        unmount_grace_period_sec = config.getint(
            CONFIG_SECTION, "unmount_grace_period_sec"
        )

        clean_up_previous_stunnel_pids()

        while True:
            config = read_config()

            check_cfs_mounts(
                config,
                child_procs,
                unmount_grace_period_sec,
                unmount_count_for_consistency,
            )
            check_child_procs(child_procs)

            time.sleep(poll_interval_sec)
    else:
        logging.info("cfs-mount-watchdog is not enabled")


if "__main__" == __name__:
    main()
