#! /usr/bin/env python3
# Copyright (C) 2020-2024 NVIDIA Corporation. ALL RIGHTS RESERVED.
# Copyright 2013,2015,2016,2017 Cumulus Networks, Inc.  All rights reserved.
# Copyright 2012 Cumulus Networks LLC, all rights reserved

#####################################################################
#
# netstat -i is a great tool for summarizing network statistics;
# however, it doesn't clear. This script saves away the
# readings at a point in time and subtracts them from the current
# readings to produce a delta. The readings are saved under the
# "/tmp/cl-netstat-<uid>" directory under a filename that is the userid
# by default. However, the user can choose to save a file with a tag
# and use that tag in future invocations to reference those readings.
#
#####################################################################

# Assume that netstat -i output looks as follows.
# Kernel Interface table
# Iface   MTU     RX-OK RX-ERR RX-DRP RX-OVR    TX-OK TX-ERR TX-DRP TX-OVR Flg
# eth0       1500     0      0      0 0             0      0      0      0 BMU
# lo        16436 246427      0      0 0        246427      0      0      0 LRU
# wlan0      1500 311162      0      0 0        196776      0      0      0 BMRU

import argparse
import subprocess
import getopt
import json
import sys
import os.path
import pickle as pickle
import re
import errno
from collections import namedtuple, OrderedDict
from tabulate import tabulate
import socket

NStats = namedtuple("NStats", "mtu, rx_ok, rx_err, rx_drop, rx_ovr, tx_ok,\
                    tx_err, tx_drop, tx_ovr, flags")
header = ['Iface', 'MTU', 'RX_OK', 'RX_ERR',
          'RX_DRP', 'RX_OVR', 'TX_OK', 'TX_ERR',
          'TX_DRP', 'TX_OVR', 'Flg']

def get_switchd_pid(pid_file='/var/run/switchd.pid'):
    try:
        with open(pid_file, 'r') as f:
            return int(f.read().strip())
    except (IOError, ValueError):
        return None


def should_discard_snapshot(cached_pid, current_pid):
    # If current_pid is None, always trust last snapshot
    if current_pid is None:
        return False
    # If switchd came up newly or restarted, discard
    if cached_pid is None or cached_pid != current_pid:
        return True
    return False


def sort_for_humans(list_foo):
    """
    Sort list_foo in the way that humans expect.
    http://nedbatchelder.com/blog/200712/human_sorting.html

    >>> sort_for_humans(['swp10', 'swp2', 'swp1', 'swp20', 'swp0', 'swp11'])
    ['swp0', 'swp1', 'swp2', 'swp10', 'swp11', 'swp20']
    """

    convert = lambda text: int(text) if text.isdigit() else text
    alphanum_key = lambda key: [convert(c) for c in re.split('([0-9]+)', key)]
    list_foo.sort(key=alphanum_key)
    return list_foo

def sort_ifaces(ifnames):
    """
    Given a list of interface names, return a list of those names sorted
    in the following order:
        lo -> eth -> swp -> bonds, bridges, etc

    >>> sort_ifaces(['bonds', 'bridges', 'swp10', 'eth0', 'swp1', 'swp20', 'swp2', 'eth1', 'lo'])
    ['lo', 'eth0', 'eth1', 'swp1', 'swp2', 'swp10', 'swp20', 'bonds', 'bridges']
    """
    sorted_variables = []
    sorted_interfaces_lo = []
    sorted_interfaces_eth = []
    sorted_interfaces_swp = []
    sorted_interfaces_nonswp = []

    for x in sort_for_humans(ifnames):
        if x.startswith('<'):
            sorted_variables.append(x)

        elif x.startswith('lo'):
            sorted_interfaces_lo.append(x)

        elif x.startswith('eth'):
            sorted_interfaces_eth.append(x)

        elif x.startswith('swp'):
            sorted_interfaces_swp.append(x)

        else:
            sorted_interfaces_nonswp.append(x)

    return (sorted_variables +
            sorted_interfaces_lo +
            sorted_interfaces_eth +
            sorted_interfaces_swp +
            sorted_interfaces_nonswp)



def cnstat_create_element(idx, netstats, procstats, ifindex):
    cntr = []
    fields = [netstats[i] for i in [1]] + \
             [procstats[i] for i in [1, 2, 3, 4, 9, 10, 11, 12]] + \
             [netstats[10]]
    cntr = NStats._make(fields)
    return dict({netstats[0]: {"ifindex": ifindex, "stats": cntr}})

def convert_int(s):
    try:
        val = int(s)
    except ValueError as v:
        val = s
    return val

def table_as_json(table):
    output = {}

    for line in table:
        if_name = line[0]

        # Build a dictionary where the if_name is the key and the value is
        # a dictionary that holds MTU, TX_DRP, etc
        output[if_name] = {
            header[1] : convert_int(line[1]),
            header[2] : convert_int(line[2]),
            header[3] : convert_int(line[3]),
            header[4] : convert_int(line[4]),
            header[5] : convert_int(line[5]),
            header[6] : convert_int(line[6]),
            header[7] : convert_int(line[7]),
            header[8] : convert_int(line[8]),
            header[9] : convert_int(line[9]),
            header[10] : line[10]
            }
    return json.dumps(output, indent=4, sort_keys=True)


def cnstat_print(cnstat_dict, use_json):
    table = []

    for if_name in sort_ifaces(list(cnstat_dict.keys())):
        entry = cnstat_dict.get(if_name)
        data = entry["stats"]
        table.append((if_name, data.mtu,
                      data.rx_ok, data.rx_err,
                      data.rx_drop, data.rx_ovr,
                      data.tx_ok, data.tx_err,
                      data.tx_drop, data.tx_err,
                      data.flags))

    if use_json:
        print(table_as_json(table))

    else:
        print('\n' + netstat_lines[0])
        print(tabulate(table, header, tablefmt='simple') + '\n')


def ns_diff(newstr, oldstr):
    new, old = int(newstr), int(oldstr)

    if new >= old:
        return str((new - old))
    else:
        return 'Unknown'

def cnstat_diff_print(cnstat_new_dict, cnstat_old_dict, use_json):
    table = []
    modified = False

    for if_name, cntr_entry in cnstat_new_dict.items():
        cntr = cntr_entry["stats"]
        current_ifindex = cntr_entry["ifindex"]
        old_entry = cnstat_old_dict.get(if_name)
        old_cntr = None
        if old_entry:
            old_cntr = old_entry.get("stats")
            old_ifindex = old_entry.get("ifindex")
            if old_ifindex != current_ifindex:
                old_cntr = None
                cnstat_old_dict.pop(if_name, None)
                modified = True

        if old_cntr is not None:
            table.append((if_name, cntr.mtu,
                          ns_diff(cntr.rx_ok, old_cntr.rx_ok),
                          ns_diff(cntr.rx_err, old_cntr.rx_err),
                          ns_diff(cntr.rx_drop, old_cntr.rx_drop),
                          ns_diff(cntr.rx_ovr, old_cntr.rx_ovr),
                          ns_diff(cntr.tx_ok, old_cntr.tx_ok),
                          ns_diff(cntr.tx_err, old_cntr.tx_err),
                          ns_diff(cntr.tx_drop, old_cntr.tx_drop),
                          ns_diff(cntr.tx_ovr, old_cntr.tx_ovr),
                          cntr.flags))
        else:
            table.append((if_name, cntr.mtu,
                          cntr.rx_ok,
                          cntr.rx_err,
                          cntr.rx_drop,
                          cntr.rx_ovr,
                          cntr.tx_ok,
                          cntr.tx_err,
                          cntr.tx_drop,
                          cntr.tx_ovr,
                          cntr.flags))

    if use_json:
        print(table_as_json(table))
    else:
        print('\n' + netstat_lines[0])
        print(tabulate(table, header, tablefmt='simple') + '\n')

    if modified:
        try:
            with open(cnstat_fqn_file, 'wb') as f:
                pickle.dump(cnstat_old_dict, f)
        except Exception as e:
            print(f"Failed to update counters file: {e}")


if __name__ == "__main__":
    parser  = argparse.ArgumentParser(description='Wrapper for netstat',
#                                      version='1.0.2',
                                      formatter_class=argparse.RawTextHelpFormatter,
                                      epilog="""
Note: Clearing stats does not affect hardware or software values.
      cl-netstat saves the current stats when -c is given, so they can
      be compared with later values.  The -c and -d options are per user (UID)
      by default.  Use the -t TAG option to change this behavior.  You must
      use the same -t TAG value with subsequent commands to get valid results.

Examples:
  cl-netstat -c -t test
  cl-netstat -t test
  cl-netstat -d -t test
  cl-netstat
  cl-netstat -r
""")
    parser.add_argument('-c', '--clear', action='store_true', help='Copy & clear stats per user (tag)')
    parser.add_argument('-d', '--delete', action='store_true', help='Delete saved stats, either the uid or the specified tag')
    parser.add_argument('-D', '--delete-all', action='store_true', help='Delete all saved stats')
    parser.add_argument('-j', '--json', action='store_true', help='Display in JSON format')
    parser.add_argument('-r', '--raw', action='store_true', help='Raw stats (unmodified output of netstat)')
    parser.add_argument('-t', '--tag', type=str, help='Save stats with name TAG', default=None)
    parser.add_argument('--clear-interface', type=str, default=None, help='clear stats for a single interface')
    args = parser.parse_args()

    save_fresh_stats = args.clear
    save_fresh_stats_single_interface = args.clear_interface
    delete_saved_stats = args.delete
    delete_all_stats = args.delete_all
    use_json = args.json
    raw_stats = args.raw
    tag_name = args.tag
    uid = str(os.getuid())
    switchd_pid = get_switchd_pid()

    if tag_name is not None:
        cnstat_file = uid + "-" + tag_name
    else:
        cnstat_file = uid

    cnstat_dir = "/tmp/cl-netstat-" + uid
    cnstat_fqn_file = cnstat_dir + "/" + cnstat_file

    if delete_all_stats:
        if os.path.exists(cnstat_dir):
            for file in os.listdir(cnstat_dir):
                os.remove(cnstat_dir + "/" + file)

            try:
                os.rmdir(cnstat_dir)
            except IOError as e:
                print(e.errno, e)
                sys.exit(e)
        sys.exit(0)

    if delete_saved_stats:
        try:
            os.remove(cnstat_fqn_file)
        except IOError as e:
            if e.errno != errno.ENOENT:
                print(e.errno, e)
                sys.exit(1)
        finally:
            if os.path.exists(cnstat_dir) and os.listdir(cnstat_dir) == []:
                os.rmdir(cnstat_dir)
            sys.exit(0)

    try:
        cmd = ['/bin/netstat', '-i']

        # If we are clearing counters get the netstat output for all interfaces
        if save_fresh_stats or save_fresh_stats_single_interface:
            cmd.append('-all')

        netstat_out = subprocess.Popen(cmd,
                                       stdout=subprocess.PIPE,
                                       shell=False).communicate()[0].decode()
    except EnvironmentError as e:
        print(e, e.errno)
        sys.exit(e.errno)

    netstat_lines = netstat_out.split("\n")

    # Since netstat -i returns some stats as 32-bits, get full 64-bit
    # stats from /prov/net/dev and display only the 64-bit stats.
    try:
        proc_out = subprocess.Popen((['/bin/cat', '/proc/net/dev']),
                                    stdout=subprocess.PIPE,
                                    shell=False).communicate()[0].decode()
    except EnvironmentError as e:
        print(e, e.errno)
        sys.exit(e.errno)

    proc = {}
    for line in proc_out.split("\n"):
        parsed = re.findall("\s*([^ ]+):(.*)", line)
        if not parsed:
            continue
        iface, stats = parsed[0]
        proc[iface] = stats.split()

    # At this point, either we'll create a file or open an existing one.
    if not os.path.exists(cnstat_dir):
        try:
            os.makedirs(cnstat_dir)
        except IOError as e:
            print(e.errno, e)
            sys.exit(1)

    # Build a dictionary of the stats
    cnstat_dict = OrderedDict()

    # Populate cnstat_dict with the current state saved to cnstat_fqn_file.
    # We will update cnstat_fqn_file with the current counters for the
    # interface we are clearing.
    if save_fresh_stats_single_interface:
        if os.path.exists(cnstat_fqn_file):
            cnstat_dict = pickle.load(open(cnstat_fqn_file, 'rb'))

    # We skip the first 2 lines since they contain no interface information
    for i in range(2, len(netstat_lines) - 1):
        netstats = netstat_lines[i].split()
        if ":" in netstats[0]:
            continue    # skip aliased interfaces

        if save_fresh_stats_single_interface is None or netstats[0] == save_fresh_stats_single_interface:
            procstats = proc.get(netstats[0])
            try:
                ifindex = socket.if_nametoindex(netstats[0])
            except OSError:
                continue
            cnstat_dict.update(cnstat_create_element(i, netstats, procstats, ifindex))

    # Now decide what information to display
    if raw_stats:
        cnstat_print(cnstat_dict, use_json)
        sys.exit(0)

    if save_fresh_stats or save_fresh_stats_single_interface:
        if switchd_pid is not None:
            cnstat_dict["_switchd_pid"] = switchd_pid
        try:
            pickle.dump(cnstat_dict, open(cnstat_fqn_file, 'wb'))
        except IOError as e:
            sys.exit(e.errno)
        else:
            if save_fresh_stats_single_interface:
                print("Cleared counters for %s" % save_fresh_stats_single_interface)
            else:
                print("Cleared counters")
            sys.exit(0)

    cnstat_cached_dict = OrderedDict()

    if os.path.isfile(cnstat_fqn_file):
        try:
            cnstat_cached_dict = pickle.load(open(cnstat_fqn_file, 'rb'))
            cached_pid = cnstat_cached_dict.pop("_switchd_pid", None)
            if should_discard_snapshot(cached_pid, switchd_pid):
                try:
                    os.remove(cnstat_fqn_file)
                except OSError:
                    pass
                cnstat_print(cnstat_dict, use_json)
            else:
                cnstat_diff_print(cnstat_dict, cnstat_cached_dict, use_json)
        except IOError as e:
            print(e.errno, e)
    else:
        if tag_name:
            print("\nFile '%s' does not exist" % cnstat_fqn_file)
            print("Did you run 'cl-netstat -c -t %s' to record the counters via tag %s?\n" % (tag_name, tag_name))
        else:
            cnstat_print(cnstat_dict, use_json)
