#!/usr/bin/env python3

import json
import onieprom
import os
import shutil
import struct
import subprocess
import sys

SYSTEM_JSON = "/run/system.json"
KKIT_IANA_PEM = 61046


class DTSystem:
    BASE = "/sys/firmware/devicetree/base"
    INFIX = BASE + "/chosen/infix"

    def __init__(self):
        self.vpdseq = 0

        dt = {}
        for root, _, files in os.walk(DTSystem.BASE):
            if "phandle" not in files:
                continue

            phandle = os.path.join(root, "phandle")
            if not os.path.exists(phandle):
                continue

            with open(phandle, "rb") as f:
                data = f.read()
            ph, = struct.unpack(">L", data)
            dt[ph] = root

        sys = {}
        for root, dirs, _ in os.walk("/sys/devices"):
            if "of_node" not in dirs:
                continue

            phandle = os.path.join(root, "of_node", "phandle")
            if not os.path.exists(phandle):
                continue

            with open(phandle, "rb") as f:
                data = f.read()
            ph, = struct.unpack(">L", data)
            if ph not in sys:
                sys[ph] = []
            sys[ph].append(root)

        phs = set(list(dt.keys()) + list(sys.keys()))

        self.devs = {ph: [Device(ph, dt.get(ph), s if s is not None else "")
                          for s in (sys.get(ph) or []) if ph is not None]
                     for ph in phs}
        self.base = Device(0, None, DTSystem.BASE)
        self.infix = Device(0, None, DTSystem.INFIX)

    def __get_phandle_array(self, name):
        path = os.path.join(DTSystem.INFIX, name)
        if not os.path.exists(path):
            return ()

        with open(path, "rb") as f:
            data = f.read()
        elems = len(data) // struct.calcsize(">L")
        return struct.unpack(">" + elems * "L", data)

    def devices_from_ph(self, ph):
        return self.devs.get(ph)

    def into_vpd(self, dev):
        def parse():
            if not dev.available():
                return {}

            try:
                with open(dev.attrpath("nvmem"), "rb", 0) as f:
                    data = onieprom.from_tlv(f)
            except:
                data = {}

            return data

        self.vpdseq += 1
        return {
            "board": dev.dtstr("infix,board", f"UNKNOWN{self.vpdseq}"),
            "available": dev.available(),
            "trusted": dev.hasdtattr("infix,trusted"),
            "data": parse(),
        }

    def infix_usb_devices(self, out):
        names = self.infix.str_array("usb-port-names", ())
        phs = self.__get_phandle_array("usb-ports")
        data = dict(zip(names, phs))
        if data != {}:
            out["usb-ports"] = []

        name_counts = {}  # Track how many times each name has been used

        for name, ph in data.items():
            for dev in self.devices_from_ph(ph):
                # Filter to only USB bus devices (usbN) that have authorized_default
                # This avoids duplicate entries for platform devices and hub interfaces
                if not dev.hasattr("authorized_default"):
                    continue

                # Convert platform path to /sys/bus/usb/devices/usbN path
                # by finding the matching symlink in /sys/bus/usb/devices/
                usb_path = dev.syspath
                try:
                    real_device = os.path.realpath(dev.syspath)
                    usb_bus_dir = "/sys/bus/usb/devices"
                    if os.path.exists(usb_bus_dir):
                        for entry in os.listdir(usb_bus_dir):
                            if not entry.startswith("usb"):
                                continue
                            candidate = os.path.join(usb_bus_dir, entry)
                            if os.path.realpath(candidate) == real_device:
                                usb_path = candidate
                                break
                except:
                    pass  # Fall back to original path if resolution fails

                # Handle name deduplication: first use gets bare name, subsequent get numbers
                if name not in name_counts:
                    name_counts[name] = 1
                    final_name = name
                else:
                    name_counts[name] += 1
                    final_name = f"{name}{name_counts[name]}"

                out["usb-ports"].append({
                    "name": final_name,
                    "path": usb_path
                })

    def infix_devices(self, kind):
        phs = self.__get_phandle_array(kind)
        return [[d for d in self.devices_from_ph(ph)] for ph in phs]

    def infix_vpds(self):
        flat_devices = [device for sublist in self.infix_devices("vpds") for device in sublist]
        return [self.into_vpd(device) for device in flat_devices]

    def vendor_name(self):
        """Extract vendor name from devicetree compatible property"""
        compatible = self.base.str_array("compatible")
        if not compatible or len(compatible) == 0:
            return None

        # Map of common devicetree vendor prefixes to proper names
        vendor_map = {
            "raspberrypi": "Raspberry Pi Foundation",
            "brcm": "Broadcom Inc.",
            "marvell": "Marvell Technology, Inc.",
            "fsl": "NXP Semiconductors N.V.",
            "nxp": "NXP Semiconductors N.V.",
            "ti": "Texas Instruments Inc.",
            "qcom": "Qualcomm Inc.",
            "rockchip": "Rockchip Electronics Co., Ltd.",
            "amlogic": "Amlogic Inc.",
            "allwinner": "Allwinner Technology Co., Ltd.",
            "mediatek": "MediaTek Inc.",
            "st": "STMicroelectronics N.V.",
            "xlnx": "Xilinx, Inc.",
            "intel": "Intel Corporation",
            "amd": "Advanced Micro Devices, Inc.",
            "nvidia": "NVIDIA Corporation",
            "bananapi": "SinoVoip Co., Ltd.",
            "sinovoip": "SinoVoip Co., Ltd.",
            "friendlyarm": "FriendlyElec",
            "friendlyelec": "FriendlyElec",
            "microchip": "Microchip Technology Inc.",
            "atmel": "Microchip Technology Inc.",
        }

        # Get the first (most specific) compatible string
        compat = compatible[0]

        # Extract vendor prefix (part before comma)
        if ',' in compat:
            prefix = compat.split(',')[0].lower()
            # Return mapped name or capitalized prefix
            return vendor_map.get(prefix, prefix.capitalize())

        return None


class DMISystem:
    BASE = "/sys/class/dmi/id"

    def read_dmi(self, attr):
        """Read DMI attribute from /sys/class/dmi/id/"""
        path = os.path.join(DMISystem.BASE, attr)
        if not os.path.exists(path):
            return None
        try:
            with open(path, 'r', encoding='utf-8') as f:
                value = f.read().strip()
                return value if value else None
        except:
            return None

    def populate(self, out):
        """Read DMI/SMBIOS data and populate output dictionary"""
        vendor = self.read_dmi("sys_vendor")
        if vendor:
            out["vendor"] = vendor

        product_name = self.read_dmi("product_name")
        if product_name:
            out["product-name"] = product_name

        serial = self.read_dmi("product_serial")
        if serial:
            out["serial-number"] = serial

        version = self.read_dmi("product_version")
        if version:
            out["product-version"] = version

    def vpds(self):
        """DMI systems don't have VPD in the traditional sense"""
        return []


class QEMUSystem:
    BASE = "/sys/firmware/qemu_fw_cfg"
    REV = BASE + "/rev"
    VPD = BASE + "/by_name/opt/vpd/raw"

    def product_vpd(self):
        data = {}
        if os.path.exists(QEMUSystem.VPD):
            try:
                with open(QEMUSystem.VPD, "rb", 0) as f:
                    data = onieprom.from_tlv(f)
            except:
                pass

        return {
            "board": "product",
            "available": os.path.exists(QEMUSystem.VPD),
            "trusted": True,
            "data": data,
        }

    def vpds(self):
        return [self.product_vpd()]

    def usb_ports(self):
        ports = [{
            "name": "USB1",
            "path": "/sys/bus/usb/devices/usb1"
        }, {
            "name": "USB2",
            "path": "/sys/bus/usb/devices/usb2"
        }]
        return ports


class Device:
    def __init__(self, ph, dtpath, syspath):
        self.ph, self.dtpath, self.syspath = ph, dtpath, syspath

    def available(self):
        return self.syspath is not None

    def __getitem__(self, attr):
        return self.attr(attr).decode("utf-8").strip("\0")

    def __setitem__(self, attr, value):
        return self.attr(attr, val=value.encode("utf-8"))

    def attrpath(self, attr):
        return os.path.join(self.syspath, attr)

    def hasattr(self, attr):
        return os.path.exists(self.attrpath(attr))

    def attr(self, attr, default=None, val=None):
        if not self.hasattr(attr):
            return default if val is None else False

        if val:
            with open(self.attrpath(attr), "wb") as f:
                f.write(val)
            return True

        with open(self.attrpath(attr), "rb") as f:
            data = f.read()
        return data

    def str(self, attr, default=None):
        val = self.attr(attr)
        return val.decode("utf-8").strip("\0") if val else default

    def str_array(self, attr, default=None):
        val = self.attr(attr)
        return val.decode("utf-8").strip("\0").split("\0") if val else default

    def dtattrpath(self, attr):
        return os.path.join(self.dtpath, attr)

    def hasdtattr(self, attr):
        return os.path.exists(self.dtattrpath(attr))

    def dtattr(self, attr, default=None):
        if not self.hasdtattr(attr):
            return default

        with open(self.dtattrpath(attr), "rb") as f:
            data = f.read()
        return data

    def dtstr(self, attr, default=None):
        val = self.dtattr(attr)
        return val.decode("utf-8").strip("\0") if val else default


def vpd_get_json_ve(vpd, pem):
    ves = vpd["data"].get("vendor-extension")
    if not ves:
        return {}

    out = {}
    for ve in filter(lambda ve: ve[0] == pem, ves):
        out.update(json.loads(ve[1]))

    return out


def vpd_get_pwhash(vpd):
    if not vpd.get("trusted"):
        return None

    kkit = vpd_get_json_ve(vpd, KKIT_IANA_PEM)
    return kkit.get("pwhash")


def vpd_inject(out, vpds):
    out["vpd"] = {vpd["board"]: vpd for vpd in vpds}

    product = out["vpd"].get("product", {}).get("data", {})
    hoistattrs = ("vendor", "product-name", "part-number", "serial-number", "mac-address")
    for attr in hoistattrs:
        if attr in product:
            out[attr] = product[attr]

    for vpd in vpds:
        pwhash = vpd_get_pwhash(vpd)
        if pwhash:
            out["factory-password-hash"] = pwhash
            break


def fallback_base_mac():
    """Find lowest valid MAC address from all interfaces."""
    base_path = '/sys/class/net'
    macs = []

    for iface in os.listdir(base_path):
        try:
            fn = os.path.join(base_path, iface, 'address')
            with open(fn, 'r', encoding='ascii') as f:
                mac = f.read().strip()

            # Skip invalid/empty/unset MAC addresses
            if not mac or len(mac) < 17 or mac == '00:00:00:00:00:00':
                continue

            # Validate MAC address format
            parts = mac.split(':')
            if len(parts) != 6:
                continue

            macs.append(mac)
        except (FileNotFoundError, ValueError):
            continue

    if macs:
        macs.sort()
        return macs[0]

    return None


def qemu_base_mac():
    """Find MAC address of first non-loopback interface, subtract with 1"""
    mac = fallback_base_mac()
    if mac:
        mac  = int(mac.replace(':', ''), 16)
        mac -= 1
        mac %= 1 << 48
        mac  = ':'.join(f"{(mac >> 8 * i) & 0xff:02x}" for i in range(5, -1, -1))
        return mac

    return None


def probe_qemusystem(out):
    """Probe Qemu based test systems and 'make run'"""
    admin_hash = "$5$mI/zpOAqZYKLC2WU$i7iPzZiIjOjrBF3NyftS9CCq8dfYwHwrmUK097Jca9A"

    qsys = QEMUSystem()
    vpds = qsys.vpds()
    usb_ports = qsys.usb_ports()
    vpd_inject(out, vpds)
    out["usb-ports"] = usb_ports
    for (attr, default) in (
            ("vendor", "QEMU"),
            ("product-name", "VM"),
            ("mac-address", qemu_base_mac()),
    ):
        if not out[attr]:
            out[attr] = default

    if os.path.exists(DMISystem.BASE):
        DMISystem().populate(out)

    if not out["factory-password-hash"] and \
       not out["vpd"]["product"]["available"]:
        # Virtual instance without VPD emulation, fallback to
        # admin/admin
        out["factory-password-hash"] = admin_hash

    # Let others react to the fact that we are running in QEMU
    subprocess.run("initctl -nbq cond set qemu".split(), check=False)
    return 0


def generic_usb_ports(out):
    """Generic USB port discovery - works for all systems.

    Each root hub gets its own uniquely-named entry so it can be independently
    enabled or disabled.  Boards that need explicit per-port control (e.g. to
    exclude internal buses) should annotate their DT with usb-ports /
    usb-port-names in the chosen/infix node instead.
    """
    ports = []
    usb_base = "/sys/bus/usb/devices"

    if not os.path.exists(usb_base):
        return

    for entry in sorted(os.listdir(usb_base)):
        if not entry.startswith("usb"):
            continue
        device_path = os.path.join(usb_base, entry)
        if not os.path.exists(os.path.join(device_path, "authorized")):
            continue
        num = entry.replace("usb", "")
        ports.append({"num": num, "path": device_path})

    if not ports:
        return

    # Single bus → "USB"; multiple → "USB1", "USB2", ... (by bus number)
    if len(ports) == 1:
        out["usb-ports"] = [{"name": "USB", "path": ports[0]["path"]}]
    else:
        out["usb-ports"] = [{"name": f"USB{p['num']}", "path": p["path"]} for p in ports]


def probe_dmisystem(out):
    """Probe DMI/SMBIOS based system (x86/AMD64)"""
    dmisys = DMISystem()

    dmisys.populate(out)
    generic_usb_ports(out)

    if not out["mac-address"]:
        out["mac-address"] = fallback_base_mac()

    vpd_inject(out, dmisys.vpds())

    return 0


def probe_dtsystem(out):
    """Probe DTS based system, expects a VPD in ONIE PROM format."""
    dtsys = DTSystem()
    vpds = dtsys.infix_vpds()

    model = dtsys.base.str("model")
    if model:
        out["product-name"] = model

    # Try devicetree-based USB discovery first, fallback to generic
    dtsys.infix_usb_devices(out)
    if "usb-ports" not in out or not out["usb-ports"]:
        generic_usb_ports(out)

    out["compatible"] = dtsys.base.str_array("compatible")

    # Extract vendor from compatible string if not already set
    if not out["vendor"]:
        vendor = dtsys.vendor_name()
        if vendor:
            out["vendor"] = vendor

    staticpw = dtsys.infix.str("factory-password-hash")
    if not out["factory-password-hash"]:
        out["factory-password-hash"] = staticpw

    vpd_inject(out, vpds)

    # Fallback to devicetree serial-number if VPD doesn't provide one
    if not out["serial-number"]:
        serial = dtsys.base.str("serial-number")
        if serial:
            out["serial-number"] = serial

    # Fallback to interface MAC if VPD doesn't provide one (e.g., SBCs)
    if not out["mac-address"]:
        out["mac-address"] = fallback_base_mac()

    return 0


def probe_wifi_radios(out):
    """Probe wifi radios via sysfs/iw and store in output dict."""
    # feature-wifi is the only Buildroot package that enables CONFIG_MAC80211/
    # CONFIG_CFG80211, so this directory exists iff WiFi support is built in.
    # feature-wifi also selects BR2_PACKAGE_IW, so iw(8) is always available
    # when this directory exists -- no separate iw check needed.
    ieee80211 = "/sys/class/ieee80211"
    if not os.path.exists(ieee80211):
        return

    radios = sorted(os.listdir(ieee80211))
    if not radios:
        return

    out["wifi-radios"] = []
    for phy in radios:
        info = {"name": phy, "bands": []}
        try:
            result = subprocess.run(
                ["iw", "phy", phy, "info"],
                capture_output=True, text=True, timeout=5
            )
            if result.returncode != 0:
                continue
        except Exception:
            continue

        freqs = []
        for line in result.stdout.splitlines():
            stripped = line.strip()
            if stripped.startswith("* ") and "MHz" in stripped:
                try:
                    freqs.append(int(float(stripped.split()[1])))
                except (ValueError, IndexError):
                    pass

        if any(2400 <= f <= 2500 for f in freqs):
            info["bands"].append({"name": "2.4GHz"})
        if any(5150 <= f <= 5900 for f in freqs):
            info["bands"].append({"name": "5GHz"})
        if any(5955 <= f <= 7115 for f in freqs):
            info["bands"].append({"name": "6GHz"})

        out["wifi-radios"].append(info)


def probe_ptp_capabilities(out):
    """Probe PTP timestamping capabilities per physical interface via ethtool --json -T.

    Only physical interfaces (those with a 'device' sysfs symlink) are probed;
    virtual interfaces such as bridges, VLANs, and tun/tap devices are skipped.
    Results are stored under out["interfaces"][<ifname>]["ptp-capabilities"].
    """
    net_base = "/sys/class/net"
    if not os.path.exists(net_base):
        return

    ifaces = {}
    for ifname in sorted(os.listdir(net_base)):
        if ifname == "lo":
            continue
        # Physical interfaces have a 'device' symlink; virtual ones do not.
        if not os.path.exists(os.path.join(net_base, ifname, "device")):
            continue

        try:
            result = subprocess.run(
                ["ethtool", "--json", "-T", ifname],
                capture_output=True, text=True, timeout=5
            )
            if result.returncode != 0:
                continue
            data = json.loads(result.stdout)[0]
        except Exception:
            continue

        caps = {
            "capabilities": data.get("capabilities", []),
            "tx-types":     data.get("tx-types", []),
            "rx-filters":   data.get("rx-filters", []),
        }

        # phc-index is -1 when no PHC is present; omit in that case.
        phc = data.get("phc-index", -1)
        if phc >= 0:
            caps["phc-index"] = phc

        # hwtstamp provider fields are present only on newer kernels/hardware.
        if (idx := data.get("hwtstamp-provider-index")) is not None:
            caps["hwtstamp-provider-index"] = idx
        if (qual := data.get("hwtstamp-provider-qualifier")) is not None:
            caps["hwtstamp-provider-qualifier"] = qual

        ifaces[ifname] = {"ptp-capabilities": caps}

    if ifaces:
        out.setdefault("interfaces", {}).update(ifaces)


def main():
    out = {
        "vendor": None,
        "product-name": None,
        "product-version": None,
        "part-number": None,
        "serial-number": None,
        "mac-address": None,
        "factory-password-hash": None,
        "vpd": {}
    }
    vpds = []

    if os.path.exists(QEMUSystem.REV):
        err = probe_qemusystem(out)
    elif os.path.exists(DMISystem.BASE):
        err = probe_dmisystem(out)
    elif os.path.exists(DTSystem.BASE):
        err = probe_dtsystem(out)
    else:
        return 1

    if err:
        return err

    probe_wifi_radios(out)
    probe_ptp_capabilities(out)

    if not out["factory-password-hash"]:
        sys.stdout.write("\n\n\033[31mCRITICAL BOOTSTRAP ERROR\n" +
                         "NO FACTORY PASSWORD FOUND\033[0m\n\n")
        err = 1

    os.umask(0o337)
    # pylint: disable=invalid-name
    with open(SYSTEM_JSON, "w", encoding="ascii") as f:
        json.dump(out, f)
    shutil.chown(SYSTEM_JSON, user="root", group="wheel")
    return err


if __name__ == "__main__":
    sys.exit(main())
