#!/opt/dosis2/bin/python3-venv
import argparse
import os
import re
from collections import defaultdict

from typing import TypeVar, Dict, List

T = TypeVar('T')


class NetworkInterface:
    def __init__(self, interface_name, device_id, hardware_address):
        # type: (str, int, str) -> None
        self.__interface_name = interface_name
        self.__device_id = device_id
        self.__hardware_address = hardware_address

    def get_interface_name(self):
        # type: () -> str
        return self.__interface_name

    def get_device_id(self):
        # type: () -> int
        return self.__device_id

    def get_hardware_address(self):
        # type: () -> str
        return self.__hardware_address

    def __repr__(self):
        # type: () -> str
        return '<%s(interface_name=%s device_id=%d hardware_address=%s)>' % (
            self.__class__.__name__,
            self.get_interface_name(),
            self.get_device_id(),
            self.get_hardware_address()
        )


def get_ethernet_interfaces():
    # type: () -> Dict[str, NetworkInterface]

    def __has_interface_files(interface_name):
        # type: (str) -> bool
        return os.path.exists(os.path.join('/sys/class/net', interface_name, 'device/device')) and \
               os.path.exists(os.path.join('/sys/class/net', interface_name, 'address'))

    interfaces = {}
    for interface in os.listdir('/sys/class/net'):
        if __has_interface_files(interface):
            with open(os.path.join('/sys/class/net', interface, 'device/device'), 'r') as f:
                device_id = int(f.read().strip(), 16)

            with open(os.path.join('/sys/class/net', interface, 'address'), 'r') as f:
                hardware_address = f.read().strip()

            interfaces[interface] = NetworkInterface(
                interface_name=interface,
                device_id=device_id,
                hardware_address=hardware_address
            )
    return interfaces


def filter_by_key(regex, data):
    # type: (str, Dict[str, T]) -> Dict[str, T]

    pattern = re.compile(regex)
    return {key: value for key, value in list(data.items()) if pattern.match(key)}


def group_by_device_id(interfaces):
    # type: (Dict[str, NetworkInterface]) -> Dict[int, List[NetworkInterface]]

    grouped = defaultdict(list)
    for interface_name, network_interface in list(interfaces.items()):
        grouped[network_interface.get_device_id()].append(network_interface)
    return grouped


def generate_persistent_network_row(name, network_interface):
    # type: (str, NetworkInterface) -> str

    return 'SUBSYSTEM=="net", ACTION=="add", DRIVERS=="?*", ATTR{address}=="%s", NAME="%s"' % (
        network_interface.get_hardware_address(),
        name
    )


def generate_persistent_network_lines(grouped_interfaces):
    # type: (Dict[int, List[NetworkInterface]]) -> List[str]

    lines = ['# This file was generated by /usr/share/dosis2/bin/generate-persistent-networks', '']

    ethnum = 0
    for device_id, interfaces in sorted(list(grouped_interfaces.items()), key=lambda keyval: len(keyval[1])):
        lines.append('# Device ID 0x%x' % (device_id,))
        for interface in sorted(interfaces, key=lambda interf: interf.get_hardware_address()):
            lines.append('%s' % (generate_persistent_network_row('eth%d' % (ethnum,), interface)))
            ethnum += 1
        lines.append('')

    return lines


def main(args):
    # type: (argparse.Namespace) -> None

    if os.environ['USER'] != 'root':
        print('this script must be ran as root')
        exit(1)
    elif not args.output:
        raise Exception('must specify an output')
    else:
        configlines = generate_persistent_network_lines(
            group_by_device_id(
                filter_by_key(
                    '^(eth|em|p\d+p\d+|enp|eno|ens).*$',
                    get_ethernet_interfaces())))

        if args.output == '-':
            print(('\n'.join(configlines)))
        else:
            with open(args.output, 'w') as f:
                f.write('\n'.join(configlines))


if __name__ == '__main__':
    __argparser = argparse.ArgumentParser(
        description='Setup deterministic network interface naming by sorting '
                    'interface count on each device, then hardware address'
    )
    __argparser.add_argument(
        '--output', '-o',
        metavar='/etc/udev/rules.d/70-persistent-net.rules',
        help='Output destination for generated persistent networks file'
    )
    main(__argparser.parse_args())
