#!/usr/bin/python3
try:
    import subprocess
    import configparser
    import errno
    import os
    import datetime
    import sys
    import hashlib
    import argparse
    import shutil
    import logging

    sys.path.append("/usr/lib/python3/site-packages/python_sdk_api")
    sys.path.append("/usr/lib/python3/dist-packages/python_sdk_api")
    from sx_api import *
    import cumulus.platforms

    import importlib.util
    import importlib.machinery

    from cumulus.mlx import mlx_open_connection
    from cumulus.mlx import mlx_close_connection
    from cumulus.mlx import mlx_invalid_port_id
    from cumulus.mlx import mlx_get_flow_control
    from cumulus.mlx import mlx_get_prio_tc_map
    from cumulus.mlx import mlx_get_port_scheduler
    from cumulus.utilities import is_bond,is_bond_member
except ImportError as e:
    pass


app_prio_switch_port_entries={}
app_prio_dict={}
roce_enabled_on_all_ports = False

# The Linux kernel populates this directory.
SYS_PATH_ROOT = "/sys/class/net"


def gen_app_prio_tlv_table_entry(app_prio_string):
    #Parse the app tlv config switchport config
    #combination and generate TLVs as per RFC
    #Octet 1 : priority 8-6 bits, sel 3-1 bits
    #Octet 2-3 : protocol id
    global app_prio_dict
    if app_prio_string in app_prio_dict.keys():
        return app_prio_dict[app_prio_string]
    else:
        try:
            app_priority_entry=""
            tlv_str_split = app_prio_string.split('_')
            proto = tlv_str_split[0]
            prio = int(tlv_str_split[2])
            prio = prio << 5
            prio_hex = hex(prio)[2:]
            port = int(tlv_str_split[1])
            port_hex = hex(port)[2:].zfill(4)
            if proto == "tcp":
                sel= 2
            elif proto == "udp":
                sel= 3
            prio_sel = prio | sel
            prio_sel_hex = hex(prio_sel)[2:].zfill(2)
            port_hex_octets = port_hex[:-2]+","+port_hex[2:]
            app_priority_entry =  prio_sel_hex + "," + port_hex_octets
            app_prio_dict[app_prio_string] = app_priority_entry
            return app_prio_dict[app_prio_string]
        except:
            return None


def add_switch_port_app_prio_entries(switchport, app_prio_table_entry):
    #Store the app tlv decoded string per port
    #this will be used to write to lldpd-auto.conf
    global app_prio_switch_port_entries
    if switchport in app_prio_switch_port_entries.keys():
        app_prio_switch_port_entries[switchport] += [app_prio_table_entry]
    else:
        app_prio_switch_port_entries[switchport] = [app_prio_table_entry]

def load_source(modname, filename):
    loader = importlib.machinery.SourceFileLoader(modname, filename)
    spec = importlib.util.spec_from_file_location(modname, filename, loader=loader)
    module = importlib.util.module_from_spec(spec)
    # The module is cached in sys.modules and not always executed
    # comment the following line to uncache the module.
    sys.modules[module.__name__] = module
    loader.exec_module(module)
    return module

def isVX():
    try:
        # Check whether we are running in VX.
        if "cumulus,vx" in subprocess.check_output('/usr/bin/platform-detect', encoding="utf8"):
            return True
    except (subprocess.CalledProcessError, OSError):
            pass

def load_pfc_config(iface):
    """
    Generate PFC map lldpcli TLV octet string for an interface from mlxcmd
    """
    port_list = mlxcmd._get_port_list_from_name_or_id(iface, False, False)
    if not port_list:
        print(('ERROR: PFC LLDP TLV enabled on invalid interface %s'%iface))
        logging.error('PFC LLDP TLV enabled on invalid interface %s'%iface)
        sys.exit(1)
    res = mlx_get_flow_control(port_list[0])
    pfc_map = 0
    pfc_dict = res['pfc']
    for k,v in list(pfc_dict.items()):
        if v['tx_enabled'] == 'yes' and v['rx_enabled'] == 'yes':
            pfc_map |= (1 << int(k))
    pfc_hex = hex(pfc_map)[2:].zfill(2)
    tlv_obj = '08,'+pfc_hex
    return tlv_obj

def load_pri_tc_map(iface):
    """
    Generate TC-SP map lldpcli TLV octet string for an interface from mlxcmd
    """
    port_list = mlxcmd._get_port_list_from_name_or_id(iface, False, False)
    if not port_list:
        print(('ERROR: ETS Config LLDP TLV enabled on invalid interface %s'%iface))
        logging.error('ETS Config LLDP TLV enabled on invalid interface %s'%iface)
        sys.exit(1)
    res = mlx_get_prio_tc_map(port_list[0])
    pri_tc = {}
    for l in res:
        if l['sp'] < 8:
            pri_tc[l['sp']] = l['tc']
        else:
            continue
    max_oct = 4
    tlv_l = []
    for i in range(max_oct):
        tlv_l.append(str(pri_tc[2*i])+str(pri_tc[2*i+1]))
    return tlv_l

def load_ets_config(iface):
    """
    Generate ETS config lldpcli TLV octet string for an interface from mlxcmd
    """
    pri_tc_l = load_pri_tc_map(iface)
    if not pri_tc_l:
        return ''
    port_list = mlxcmd._get_port_list_from_name_or_id(iface, False, False)
    res = mlx_get_port_scheduler(port_list[0])

    sched_j_sg = {}
    for el in res:
        if el['level'] == 2 and el['index'] < 8:
            sd = sched_j_sg[el['index']] = {}
            sd['sta'] = el['dwrr_mode']
            sd['bw_pc'] = el['dwrr_weight']
    bw_pl = []
    sta_l = []
    sta_d = {'strict priority': 0, 'dwrr': 2}
    for k,v in list(sched_j_sg.items()):
        bw_pl.append(hex(v['bw_pc'])[2:].zfill(2))
        sta_l.append(hex(sta_d[v['sta']])[2:].zfill(2))
    ets_oui_i = pri_tc_l + bw_pl + sta_l
    ets_cfg_s = '00,'+','.join(ets_oui_i)
    return ets_cfg_s


def gen_lldp_dcbx_tlv(obj , swp, tlv_type):
    """
    Generate lldpcli string for tlv_type and swp
    """
    global roce_enabled_on_all_ports
    glob_conf ='configure'
    lldp_cnf_str = ''
    if not obj:
        return lldp_cnf_str
    base = ['configure ports ',' lldp custom-tlv ', 'replace ', 'oui 00,80,c2 subtype ', ' oui-info ']
    if swp == 'allPorts':
        if tlv_type == 'pfc':
            lldp_cnf_str = glob_conf + base[1] + base[3] + '11' + base[4] + obj
        if tlv_type == 'ets-config':
            lldp_cnf_str = glob_conf + base[1] + base[3] + '9' + base[4] + obj
        if tlv_type == 'ets-recomm':
            lldp_cnf_str = glob_conf + base[1] + base[3] + '10' + base[4] + obj
        if tlv_type == 'roce-app':
            lldp_cnf_str = glob_conf + base[1] + base[3] + '12' + base[4] + obj
    else:
        if tlv_type == 'pfc':
            lldp_cnf_str = base[0] + swp + base[1] + base[3] + '11' + base[4] + obj
        if tlv_type == 'ets-config':
            lldp_cnf_str = base[0] + swp + base[1] + base[3] + '9' + base[4] + obj
        if tlv_type == 'ets-recomm':
            lldp_cnf_str = base[0] + swp + base[1] + base[3] + '10' + base[4] + obj
        if tlv_type == 'app-prio-tlv':
            if roce_enabled_on_all_ports:
                lldp_cnf_str = base[0] + swp + base[1] + base[2] + base[3] + '12' + base[4] + obj
            else:
                lldp_cnf_str = base[0] + swp + base[1] + base[3] + '12' + base[4] + obj
    return lldp_cnf_str


def get_active_swp_iface():
    """
    Returns an active swp interface
    """
    linux_to_sdk = mlxcmd.get_linux_to_sdk()
    for k in linux_to_sdk.keys():
        if (not is_bond(k)):
            return k

def get_active_non_bond_mem_iface():
    """
    Returns an active bond interface or swp interface, which isn't a bond-member
    """
    linux_to_sdk = mlxcmd.get_linux_to_sdk()
    for k in linux_to_sdk.keys():
        if (not is_bond_member(k)):
            return k

def lldp_gen_pfc_tlv(iface):
    """
    Generate PFC MAP lldpcli string for an interface
    """
    if iface != 'allPorts':
        pfc_map = load_pfc_config(iface)
    else:
        pfc_map = load_pfc_config(get_active_swp_iface())
    pfc_tlv_str = gen_lldp_dcbx_tlv(pfc_map, iface, tlv_type = 'pfc')
    return pfc_tlv_str

def lldp_gen_ets_cfg_tlv(iface):
    """
    Generate ETS config lldpcli string for an interface
    """
    if iface != 'allPorts':
        ets_cfg = load_ets_config(iface)
    else:
        ets_cfg = load_ets_config(get_active_non_bond_mem_iface())
    ets_cfg_tlv_str = gen_lldp_dcbx_tlv(ets_cfg, iface, tlv_type = 'ets-config')
    return ets_cfg_tlv_str

def lldp_gen_ets_rec_tlv(iface):
    """
    Generate ETS recommendation lldpcli string for an interface
    """
    if iface != 'allPorts':
        ets_rec = load_ets_config(iface)
    else:
        ets_rec = load_ets_config(get_active_non_bond_mem_iface())
    ets_rec_tlv_str = gen_lldp_dcbx_tlv(ets_rec, iface, tlv_type = 'ets-recomm')
    return ets_rec_tlv_str

def lldp_gen_app_prio_tlv(iface, total_app_tlvs):
    """
    Generate DCBX app prio lldpcli string for an interface
    """
    # prio -3 Sel -3 , protocol ID- 4791 (0x12b7)
    # 1st octet- reserved, 2nd oct - prio (1-3 bits) reserved (4-5 bits) Sel (6-8 bits)
    # 3rd & 4th oct - protocol ID
    app_prio = ",".join(app_prio_switch_port_entries[iface])
    app_tlv_str = gen_lldp_dcbx_tlv("00,"+total_app_tlvs, iface, tlv_type = 'app-prio-tlv')
    return app_tlv_str

def lldp_gen_roce_app_prio_tlv(iface):
    """
    Generate ROCE app prio lldpcli string for an interface
    """
    # prio -3 Sel -3 , protocol ID- 4791 (0x12b7)
    # 1st octet- reserved, 2nd oct - prio (1-3 bits) reserved (4-5 bits) Sel (6-8 bits)
    # 3rd & 4th oct - protocol ID
    roce_app_prio = '00,63,12,b7'
    roce_app_tlv_str = gen_lldp_dcbx_tlv(roce_app_prio, iface, tlv_type = 'roce-app')
    return roce_app_tlv_str

def generate_hash(file_name, block = 8192):
    """
    Read file in blocks and generate md5 hash of the file
    """
    m = hashlib.md5()
    with open(file_name, "rb") as f:
        while True:
            buf = f.read(block)
            if not buf:
                break
            m.update(buf)
    return m.hexdigest()

if __name__ == '__main__':
    restart_lldpd = False

    if isVX():
        exit(0)

    try:
        mlxcmd = load_source('mlxcmd','/usr/lib/cumulus/mlxcmd')
    except Exception as e:
        logging.error('failed to load mlxcmd: %s'%e)
        sys.exit(1)

    parser = argparse.ArgumentParser()
    parser.add_argument('-r', '--restart_lldpd', action="store_true", help="Restart lldpd")

    args = parser.parse_args()
    if args.restart_lldpd:
        restart_lldpd = True
    config = configparser.ConfigParser()
    config.read('/etc/cumulus/lldp-dcbx-nvue.conf')
    pfc_en_ifaces = config['lldp']['interfaces.pfc_tlv_enable'].split(',')
    ets_cfg_en_ifaces = config['lldp']['interfaces.ets_cfg_tlv_enable'].split(',')
    ets_rec_en_ifaces = config['lldp']['interfaces.ets_rec_tlv_enable'].split(',')
    roce_en_ifaces = config['lldp']['interfaces.roce_app_tlv_enable'].split(',')
    roce_enabled_on_all_ports = False
    if 'allPorts' in roce_en_ifaces:
        roce_enabled_on_all_ports = True

    try:
        app_tlv_cfg_list = config['lldp']['interfaces.app_priority_tlv_configs'].split(',')
    except:
        app_tlv_cfg_list = []

    if restart_lldpd:
        f = open('/tmp/lldpd-auto.conf', 'w')
    else:
        f = open('/etc/lldpd.d/lldpd-auto.conf', 'w')

    mlx_open_connection()
    f.write('\n############ PFC MAP TLV ################\n\n')
    for iface in pfc_en_ifaces:
        if not iface:
            continue
        pfc = lldp_gen_pfc_tlv(iface)
        f.write(pfc + '\n')

    f.write('\n############ ETS CFG TLV ################\n\n')
    for iface in ets_cfg_en_ifaces:
        if not iface:
            continue
        f.write(lldp_gen_ets_cfg_tlv(iface) + '\n')

    f.write('\n############ ETS REC TLV ################\n\n')
    for iface in ets_rec_en_ifaces:
        if not iface:
            continue
        f.write(lldp_gen_ets_rec_tlv(iface) + '\n')

    f.write('\n############ ROCE APP TLV ################\n\n')
    for iface in roce_en_ifaces:
        if not iface:
            continue
        f.write(lldp_gen_roce_app_prio_tlv(iface) + '\n')

    f.write('\n############ APP PRIO TLV ################\n\n')
    for entry in app_tlv_cfg_list:
        if not entry:
            continue
        switchport , prio_string = entry.split('_',1)
        table_entry = gen_app_prio_tlv_table_entry(prio_string)
        if table_entry:
            add_switch_port_app_prio_entries(switchport,table_entry)

    for switchport in app_prio_switch_port_entries.keys():
        if roce_enabled_on_all_ports:
            total_tlv_entries = "63,12,b7,"+",".join(app_prio_switch_port_entries[switchport])
        else:
            total_tlv_entries = ",".join(app_prio_switch_port_entries[switchport])
        app_tlv_str = lldp_gen_app_prio_tlv(switchport, total_tlv_entries)
        f.write(app_tlv_str + '\n')

    mlx_close_connection()
    f.close()
    if restart_lldpd:
        hash_o = '-1'
        hash_n = generate_hash("/tmp/lldpd-auto.conf")
        if os.path.isfile('/etc/lldpd.d/lldpd-auto.conf'):
            hash_o = generate_hash("/etc/lldpd.d/lldpd-auto.conf")
        if hash_n != hash_o:
            shutil.move(os.path.join('/tmp/', 'lldpd-auto.conf'),
                        os.path.join('/etc/lldpd.d/', 'lldpd-auto.conf'))
            out = subprocess.check_output(['systemctl', 'reset-failed', 'lldpd'])
            if out:
                logging.error(out)
            out = subprocess.check_output(['systemctl', 'restart', 'lldpd'])
            if out:
                logging.error(out)
