#!/usr/bin/env python3
# Copyright (c) 2022, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# NVIDIA CORPORATION and its licensors retain all intellectual property
# and proprietary rights in and to this material, related documentation
# and any modifications thereto. Any use, reproduction, disclosure or
# distribution of this material and related documentation without an express
# license agreement from NVIDIA CORPORATION is strictly prohibited.
#
# SPDX-FileCopyrightText: Copyright (c) 2022 NVIDIA CORPORATION & AFFILIATES
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary

import sys
import json
import click
import subprocess

from tabulate import tabulate
from natsort import natsorted
from click_default_group import DefaultGroup

MLNX_ID = "15b3"
SPC1_PRODUCT_ID = "cb84"
CONTEXT_SETTINGS = dict(help_option_names=['-h', '--help', '-?'])
WHAT_JUST_HAPPENED_SERVICE_NAME = "what-just-happened"
WHAT_JUST_HAPPENED_CONFIG_FILE_PATH = "/etc/what-just-happened/what-just-happened.json"


def is_wjh_active():
    return subprocess.call(
        ["systemctl", "is-active", "what-just-happened.service"],
        stdout=subprocess.DEVNULL,
        stderr=subprocess.DEVNULL
    ) == 0


def is_spc1():
    cmd = 'lspci -n | grep "15b3:cb84"'.format(MLNX_ID, SPC1_PRODUCT_ID)
    try:
        proc = subprocess.Popen(cmd, stdout=subprocess.PIPE, shell=True, text=True)
        stdout = proc.communicate()[0]
        proc.wait()
    except OSError as e:
        raise OSError("Could not execute 'lspci' command: %s" % str(e))
    if stdout != "":
        return True
    return False


def is_wjh_docker_up():
    cmd = "docker ps | grep {}".format(WHAT_JUST_HAPPENED_SERVICE_NAME)
    try:
        proc = subprocess.Popen(cmd, stdout=subprocess.PIPE, shell=True, text=True)
        stdout = proc.communicate()[0]
        proc.wait()
    except OSError as e:
        raise OSError("Could not execute 'docker ps' command: %s" % str(e))
    if stdout != "":
        return True
    return False


def get_channels():
    try:
        return json.loads(
            "\n".join([line for line in subprocess.check_output([
                "/bin/bash", "-c",
                "docker exec -i {} /usr/bin/wjhcli -g getchannels".format(
                    WHAT_JUST_HAPPENED_SERVICE_NAME)]
            ).decode().splitlines() if line[0] in ["{", " ", "}", "\t"]])
        )
    except Exception as e:
        print(str(e))
        return {}


#
# 'cli' group (root group)
#

# This is our entrypoint - the main "show" command
@click.group(
    default="poll",
    cls=DefaultGroup,
    default_if_no_args=True,
    context_settings=CONTEXT_SETTINGS
)
@click.pass_context
def cli(ctx):
    """NVIDIA Cumulus Linux - what-just-happened command line"""
    if ctx.invoked_subcommand is None:
        ctx.forward(poll)


@cli.command()
def dump():
    """Dump what-just-happened debug information"""
    if not is_wjh_active():
        click.echo(
            '"{}" feature is not active, run "systemctl start {}"'
            .format(WHAT_JUST_HAPPENED_SERVICE_NAME, WHAT_JUST_HAPPENED_SERVICE_NAME),
            err=True
        )
        return
    cmd = "docker exec -i {} /usr/bin/wjhcli -d".format(WHAT_JUST_HAPPENED_SERVICE_NAME)
    proc = subprocess.Popen(["/bin/bash", "-c", cmd], stdout=sys.stdout, stderr=sys.stderr)
    proc.wait()


@cli.command()
@click.argument('channels', nargs=-1)
@click.option('--aggregate', is_flag=True, help='Dump aggregated counters')
@click.option('--export', is_flag=True, help='Save droped packets into pcap file')
@click.option('--no_metadata', is_flag=True, help='Save the pcap file without metadata')
def poll(channels, aggregate, export, no_metadata):
    """Poll what-just-happened user channel"""

    if no_metadata and not export:
        click.echo("'--no_metadata' option is available only with '--export'. \nAborting!", err=True)
        return

    if not is_wjh_active():
        click.echo(
            '"{}" feature is disabled, run "systemctl start {}"'
            .format(WHAT_JUST_HAPPENED_SERVICE_NAME, WHAT_JUST_HAPPENED_SERVICE_NAME),
            err=True)
        return

    if not is_wjh_docker_up():
        click.echo('"docker" service is not running. Please retry again in few seconds', err=True)
        return

    if not channels:
        channels = get_channels().keys()
        if not channels:
            click.echo("No what-just-happened channels configured")
            return
        if aggregate and len(channels) > 1 and ("layer-1" in channels):
            channels.remove("layer-1")
        if is_spc1() and ("buffer" in channels):
            channels.remove("buffer")

    # Different channels have different aggregate tables, so if aggregated flag is raised
    # and there is more then one channels, abort.
    if aggregate and len(channels) > 1:
        err_msg = "{} channels have different aggregate tables.\nAborting!".format(
            ", ".join([chan for chan in channels])
        )
        click.echo(err_msg, err=True)
        return

    cmd = "docker exec -i {} /usr/bin/wjhcli {} {} {} {}".format(
        WHAT_JUST_HAPPENED_SERVICE_NAME,
        " ".join(["-c {}".format(chan) for chan in channels]),
        "-p" if export else "",
        "-a" if aggregate else "",
        "-m" if no_metadata else ""
    )
    proc = subprocess.Popen(["/bin/bash", "-c", cmd], stdout=sys.stdout, stderr=sys.stderr)
    proc.wait()


@cli.group()
def configuration():
    """what-just-happened configuration"""
    pass


@configuration.command()
def channels():
    """ what-just-happened channel configuration  """
    wjh_channel_table = get_channels()
    header = ['Channel', 'Type', 'Drop Groups']
    body = []
    for channel in natsorted(wjh_channel_table.keys()):
        channel_config = {k.lower(): v for k, v in wjh_channel_table[channel].items()}

        drop_category_list = channel_config.get('drop_category_list', [])

        if not drop_category_list:
            drop_category = "N/A"
        else:
            drop_category = ", ".join(drop_category_list)

        body.append([
            channel,
            channel_config.get('type', 'N/A'),
            drop_category,
        ])
    click.echo(tabulate(body, header))


if __name__ == "__main__":
    cli()
