#! /usr/bin/env python3

# SPDX-FileCopyrightText: Copyright (c) 2022-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# NVIDIA CORPORATION, its affiliates and 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 or
# its affiliates is strictly prohibited.

from enum import Enum
import logging
import os
import re
import sys
from time import sleep
import bmc_interface
from argparse import ArgumentParser

# Configure logging - only for our logger, not root logger
logger = logging.getLogger('TB500CTL')
logger.setLevel(logging.INFO)

# Create console handler with formatting
console_handler = logging.StreamHandler()
formatter = logging.Formatter(
    '%(name)s: %(levelname)s: %(message)s'
)
console_handler.setFormatter(formatter)
logger.addHandler(console_handler)

try:
    from supported_targets_all import supported_targets
except:
    from supported_targets import supported_targets

logger.info("version 20251204.0")

class BootDev(Enum):
    QSPI0 = 0x0
    # rsvd
    # rsvd
    UART0 = 0x3
    OOBHUB = 0x4
    RCM1 = 0x5 # not supported
    RCM2 = 0x6

class RcmDev(Enum):
    OCP_RC = 0x0
    RCM1 = 0x1 # not supported
    RCM2 = 0x2

class UartDev(Enum):
    UART_BRIDGE = 0 # UART0 - UART1
    UART_HDR = 1    # UARTn - pin header on board
    MCU = 2         # UARTn - MCU
    USB = 3         # UARTn - USB

class QspiDev(Enum):
    QSPI_BRIDGE = 0 # QSPI0 - QSPI1
    FLASH = 1       # QSPIn - QSPIn_FLASH
    OOBHUB = 2      # OOBHUB - QSPIn_FLASH
    SPI1 = 3        # QSPI1 - SPI1 (only on QSPI1)

class UsbMgmt0Dev(Enum):
    MGMT0 = 0
    MGMT1_HUB = 1
    J2002_CON = 2

class UsbMctpDev(Enum):
    HOST = 0
    BMC = 1

class Usb2Dev(Enum):
    HOST = 0
    BMC = 1

# NOTE: any value over 0 is secondary with two bits available.
#       we only do up to c2 config so just doing this for simplification.
class PackageId(Enum):
    PRIMARY = 0
    SECONDARY = 1

r'''
E5010/E5020:

    i2c_mgmt0_brd0 ---- i2c_bmc_mgmt0 ----- 0x20 i2c_mgmt0
                                    | |- 0x21 i2c_mgmt0
                                    |
                                    |--- 0x22 i2c mux ctl (uart mux)
                                        |- 0x23 i2c mux ctl
                                        |- 0x24 i2c mux ctl

    i2c_hmc_mgmt1 ----- HMC
                MUX \-- HMC
                    MUX \-- i2c_bmc_mgmt1 ----- 0x20 i2c_mgmt1
                                            |- 0x21 i2g_mgmt1 (straps)
'''

STRAPS_0      = "0x20"
STRAPS_1      = "0x21"
UART_MUX_CTL  = "0x22"
USB_MUX_CTL   = "0x23"
OOB_CTL       = "0x51"

socket_list = [
    {
        "mgmt0": 2,
        "mgmt1": 1,
        "mcu": 70,
    },
    {
        "mgmt0": 6,
        "mgmt1": 7,
        "mcu": 71,
    }
]

e5010_features = ["uart_mux", "usb_mgmt_mux", "usb2_mux", "qspi_mux", "package_id"]
p5035_features = ["usb_mctp_mux"] # NOTE: A1 rework allows usb mux, A2 will include by default
pg558_features = ["hmc", "multi_board"]

bmc = None

def main():
    global socket_list
    global bmc

    tb500_targets = supported_targets["tb500"]
    commands = "status | defaults | enter_ist | power_{on,off} | power_cycle | os_restart | warm_reset | usbrcm_reset | streaming_reset" \
    "| ist_{up,down} | rcm_{up,down} | jtag_sel_{on,off} | boot_sel_{qspi,uart,oobhub,usb} | rcm_sel_{ocp_rc,usb} " \
    "| uart{0,1}_sel_{host,bridge,hdr,mcu,usb} | usb_mgmt_sel_{mgmt0,mgmt1,con} | usb_mctp_sel_{con,mgmt1}" \
    "| usb2_sel_{host,bmc} | package_sel_{primary,secondary} | die_sel_{##} | qspi{0,1}_sel_{bridge,flash,oobhub,spi1}"

    parser = ArgumentParser(epilog="Order: straps -> reset -> muxes.\nAvailable commands:  " + commands,
            usage="tb500-ctl [options] <comma delimited command list>")
    parser.set_defaults()
    parser.add_argument("--target", "-t", action="store", type=str,
                      dest="target", help="Target board [%s]" % " | ".join(tb500_targets))
    parser.add_argument("--variant", "-v", action="store", type=str,
                      dest="variant", help="Target board variant [A00 | A01 | A02 | A03 | ...]")
    parser.add_argument("--power_delay", action="store", type=int,
                      dest="power_delay", help="delay after reset time in seconds",
                      default=30)
    parser.add_argument("--i2c_delay", action="store", type=float,
                      dest="i2c_delay", help="delay between I2C commands, helpful for fw issues",
                      default=0)
    parser.add_argument("--bmc", "-b", action="store", type=str,
                      dest="bmc", help="IP address or hostname of target BMC. Used for SSH and Redfish. Can also use 'BMC_URL' envvar.",
                      default=None)
    parser.add_argument("--username", "-u", action="store", type=str,
                      dest="username", help="BMC username for Redfish and SSH login",
                      default="root")
    parser.add_argument("--password", "-p", action="store", type=str,
                      dest="passwd", help="BMC password for Redfish and SSH login",
                      default="0penBmc")
    parser.add_argument("--socket", "-s", action="append", type=int, dest="sockets", metavar="SOCKET",
                      help="Specifies socket index to add for current configuration (one -s per socket).")
    parser.add_argument("--debug", action="store_true",
                      help="Enable debug output of this script.",
                      default=False)
    parser.add_argument("--concise", action="store_true",
                      help="Enable concise output of this script.",
                      default=False)
    parser.add_argument("--skip_mux", action="store_true",
                      help="Do not touch muxes, used for improper fw or mux cfg.",
                      default=False)
    parser.add_argument("--skip_lock", action="store_true",
                      help="Do not touch or check lock file (i.e. for GVS).",
                      default=False)

    (options, args) = parser.parse_known_args()

    if options.debug:
        logger.setLevel(logging.DEBUG)
    elif options.concise:
        logger.setLevel(logging.WARNING)

    if not options.target:
        logger.error("No target specified")
        sys.exit(-1)

    options.target = options.target.upper()

    socket_cnt = 1
    if not options.sockets:
        options.sockets = [0]
    else:
        socket_cnt = len(options.sockets)

    features = None
    mgmt0_ports = []
    mgmt1_ports = []
    if "E5010" in options.target or "E5020" in options.target:
        features = e5010_features
        mgmt0_ports = [STRAPS_0, STRAPS_1, UART_MUX_CTL, USB_MUX_CTL]
        mgmt1_ports = [STRAPS_0, STRAPS_1]
    elif "P5035" in options.target:
        features = p5035_features
        mgmt0_ports = [STRAPS_0, STRAPS_1]
        mgmt1_ports = [STRAPS_0, STRAPS_1]
    elif "PG558" in options.target:
        features = pg558_features
        mgmt1_ports = [STRAPS_0, STRAPS_1]
    else:
        logger.error("Invalid target")
        sys.exit(-1)

    if "multi_board" not in features and socket_cnt != 1:
        logger.error("Multi-board multisocket not supported on this target!")
        sys.exit(-1)

    if socket_cnt > 2:
        logger.error("Only 1-2 sockets supported!")
        sys.exit(-1)

    if socket_cnt != len(set(options.sockets)):
        logger.error("socket id(s) specified more than once")
        sys.exit(-1)

    try:
        socket_list = [socket_list[i] for i in options.sockets]
    except:
        logger.error("invalid socket list specified")
        sys.exit(-1)

    if options.bmc is None:
        if "BMC_URL" in os.environ:
            options.bmc = re.sub(r'^https?://', '', os.environ["BMC_URL"])
            logger.info(f"Using --bmc={options.bmc} from 'BMC_URL' envvar")
        elif "BMCURL" in os.environ:
            options.bmc = re.sub(r'^https?://', '', os.environ["BMCURL"])
            logger.info(f"Using --bmc={options.bmc} from 'BMCURL' envvar")
        elif "BMCIP" in os.environ:
            options.bmc = os.environ["BMCIP"]
            logger.info(f"Using --bmc={options.bmc} from 'BMCIP' envvar")
        else:
            logger.error("bmc functionality is required for TB500 and needs --bmc or BMC_URL envvar")
            sys.exit(-1)

    bmc = bmc_interface.bmc(host=options.bmc, features=features, logger=logger, user=options.username, passwd=options.passwd,
        i2c_delay=options.i2c_delay, skip_lock=options.skip_lock)

    try:
        i2c_enable(mgmt0_ports, mgmt1_ports)
    except:
        logger.error("Failed to enable/find required I2C interfaces")
        bmc.close()
        sys.exit(-1)

    logger.info("Initialized I2C management interface")

    if len(args) < 1:
        logger.error(f"Board control command missing.  Must be one of: {commands}")
        bmc.close()
        sys.exit(-1)

    args = args[0].split(",")

    # command parse
    for arg in args:
        if arg == "status":
            print(bmc.info(options.concise))
            gpio_info(options.concise, options.skip_mux)
        elif arg == "defaults":
            set_defaults(skip_mux=options.skip_mux)
        elif arg == "rcm_power_on":
            logger.warning("LEGACY! please use usbrcm_reset or streaming_reset")
            set_defaults(skip_mux=options.skip_mux)
            pull_rcm(0) # RCM to 0 enables RCM boot
            bmc.power_reset("ForceOn")
            sleep(options.power_delay)
        elif arg == "usbrcm_reset":
            boot_sel(BootDev.QSPI0)
            rcm_sel(RcmDev.RCM2)
            pull_rcm(0) # FORCED_RECOVERY on
            bmc.power_reset("PowerCycle")
            sleep(options.power_delay)
        elif arg == "streaming_reset":
            boot_sel(BootDev.RCM2)
            rcm_sel(RcmDev.OCP_RC)
            pull_rcm(1) # FORCED_RECOVERY off
            bmc.power_reset("PowerCycle")
            sleep(options.power_delay)
        elif arg == "enter_ist":
            set_defaults(skip_mux=options.skip_mux)
            pull_ist(1) # IST to 1 enables IST boot
            bmc.power_reset("ForceRestart")
            sleep(options.power_delay)
        elif arg == "power_on":  # for cold boot
            bmc.power_reset("On")
            sleep(options.power_delay)
        elif arg == "power_off":
            bmc.power_reset("GracefulShutdown")
            sleep(options.power_delay)
        elif arg == "force_on" or arg == "force_off" or arg == "force_reset":
            logger.error("Power forcing commands are not supported on TB500")
            bmc.close()
            sys.exit(-1)
        elif arg == "os_restart": # for complete system reset (graceful)
            bmc.power_reset("GracefulRestart")
            sleep(options.power_delay)
        elif arg == "power_cycle":  # kills power rails
            bmc.power_reset("PowerCycle")
            sleep(options.power_delay)
        elif arg == "warm_reset":  # cpu reset only
            bmc.power_reset("ForceRestart")
            sleep(options.power_delay)
        elif arg == "ist_down":
            pull_ist(1)
        elif arg == "ist_up":
            pull_ist(0)
        elif arg == "rcm_down":
            pull_rcm(0)
        elif arg == "rcm_up":
            pull_rcm(1)
        elif arg == "boot_sel_uart":
            boot_sel(BootDev.UART0)
        elif arg == "boot_sel_qspi":
            boot_sel(BootDev.QSPI0)
        elif arg == "boot_sel_oobhub":
            boot_sel(BootDev.OOBHUB)
        elif arg == "boot_sel_usb":
            boot_sel(BootDev.RCM2)
        elif arg == "rcm_sel_ocp_rc":
            rcm_sel(RcmDev.OCP_RC)
        elif arg == "rcm_sel_usb":
            rcm_sel(RcmDev.RCM2)
        elif arg == "jtag_sel_off":
            jtag_sel(0)
        elif arg == "jtag_sel_on":
            jtag_sel(1)
        elif arg.startswith("die_sel_"):
            die_sel(arg)
        elif arg == "uart0_sel_host":
            logger.error("uart0_sel_host is deprecated--please use external UART-to-USB or apply manually.")
        elif arg == "uart0_sel_bridge":
            uart_sel(0, UartDev.UART_BRIDGE)
        elif arg == "uart0_sel_hdr":
            uart_sel(0, UartDev.UART_HDR)
        elif arg == "uart0_sel_mcu":
            uart_sel(0, UartDev.MCU)
        elif arg == "uart0_sel_usb":
            uart_sel(0, UartDev.USB)
        elif arg == "uart1_sel_host":
            logger.error("uart1_sel_host is deprecated--please use external UART-to-USB or apply manually.")
        elif arg == "uart1_sel_bridge":
            uart_sel(1, UartDev.UART_BRIDGE)
        elif arg == "uart1_sel_hdr":
            uart_sel(1, UartDev.UART_HDR)
        elif arg == "uart1_sel_mcu":
            uart_sel(1, UartDev.MCU)
        elif arg == "uart1_sel_usb":
            uart_sel(1, UartDev.USB)
        elif arg == "usb_mgmt_sel_mgmt0":
            usb_mgmt_sel(UsbMgmt0Dev.MGMT0)
        elif arg == "usb_mgmt_sel_mgmt1":
            usb_mgmt_sel(UsbMgmt0Dev.MGMT1_HUB)
        elif arg == "usb_mgmt_sel_con":
            usb_mgmt_sel(UsbMgmt0Dev.J2002_CON)
        elif arg == "usb_mctp_sel_con" or "usb_mctp_sel_host":
            usb_mctp_sel(UsbMctpDev.HOST)
        elif arg == "usb_mctp_sel_mgmt1" or "usb_mctp_sel_host":
            usb_mctp_sel(UsbMctpDev.BMC)
        elif arg == "usb2_sel_d1" or arg == "usb2_sel_host":
            usb2_sel(Usb2Dev.HOST)
        elif arg == "usb2_sel_d2" or arg == "usb2_sel_bmc":
            usb2_sel(Usb2Dev.BMC)
        elif arg == "package_sel_primary":
            package_id_sel(PackageId.PRIMARY)
        elif arg == "package_sel_secondary":
            package_id_sel(PackageId.SECONDARY)
        elif arg == "qspi0_sel_bridge":
            qspi_sel(0, QspiDev.QSPI_BRIDGE)
        elif arg == "qspi0_sel_flash":
            qspi_sel(0, QspiDev.FLASH)
        elif arg == "qspi0_sel_oobhub":
            qspi_sel(0, QspiDev.OOBHUB)
        elif arg == "qspi1_sel_bridge":
            qspi_sel(1, QspiDev.QSPI_BRIDGE)
        elif arg == "qspi1_sel_flash":
            qspi_sel(1, QspiDev.FLASH)
        elif arg == "qspi1_sel_oobhub":
            qspi_sel(1, QspiDev.OOBHUB)
        elif arg == "qspi1_sel_spi1":
            qspi_sel(1, QspiDev.SPI1)
        else:
            logger.error(f"Unknown board control command '{arg}'.  Must be one of: {commands}")
            bmc.close()
            sys.exit(-1)

    bmc.close()

def gpio_prefix(name, brd, module):
    prefix = f"B{brd}_M{module}_"
    if bmc.has_feature("hmc"):
        prefix = f"BRD{brd}_"
    return f"{prefix}{name}"

def i2c_enable(mgmt0_ports, mgmt1_ports):
    # if not using HMC, configure i2c mux to BMC
    if not bmc.has_feature("hmc"):
        bmc.tca95xx_write(5, STRAPS_0, 0, 5, 0)
        bmc.tca95xx_write(5, STRAPS_1, 0, 0, 0)

    for socket in socket_list:
        if mgmt0_ports != [] and not bmc.check_file(f"/dev/i2c-{socket["mgmt0"]}"):
            logger.error("Could not find I2C MGMT0 bus!")
            bmc.close()
            sys.exit(-1)
        if mgmt1_ports != [] and not bmc.check_file(f"/dev/i2c-{socket["mgmt1"]}"):
            logger.error("Could not find I2C MGMT1 bus!")
            bmc.close()
            sys.exit(-1)

        if bmc.has_feature("qspi_mux") and not bmc.check_file(f"/dev/i2c-{socket["mcu"]}"):
            logger.warning("MCU not connected, skipping QSPI mux")
            bmc.remove_feature("qspi_mux")

# set force recovery strap and reset
def pull_rcm(val):
    logger.info(f"Holding FORCED_RECOVERY_L {'high' if val else 'low'} via BMC...")
    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("CPU_FORCED_RECOVERY_L-B", socket_idx, 1), val)

# set ist strap and reset
def pull_ist(val):
    logger.info(f"Holding IST_BOOT {'high' if val else 'low'} via BMC...")
    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("CPU_IST_BOOT-B", socket_idx, 1), val)

# set jtag strap
def jtag_sel(val):
    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("NVJTAG_SEL_S-B", socket_idx, 1), val)
    logger.info("Enabled NVJTAG for selected sockets")

# set boot device selection strap
def boot_sel(id=BootDev.QSPI0):
    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("CPU_BOOT_DEV_SEL0-B", socket_idx, 1), id.value & 0x1)
        bmc.gpio_set(gpio_prefix("CPU_BOOT_DEV_SEL0-B", socket_idx, 1), (id.value >> 1) & 0x1)
        bmc.gpio_set(gpio_prefix("CPU_BOOT_DEV_SEL0-B", socket_idx, 1), (id.value >> 2) & 0x1)
    logger.info(f"Selected boot dev {id.name} for next boot")

# set recovery device strap
def rcm_sel(id=RcmDev.RCM2):
    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("CPU_RECOVERY_TYPE0-B", socket_idx, 1), id.value & 0x1)
        bmc.gpio_set(gpio_prefix("CPU_RECOVERY_TYPE0-B", socket_idx, 1), (id.value >> 1) & 0x1)
    logger.info(f"Selected RCM dev {id.name} for next boot")

# set uart mux configuration
def uart_sel(uart, id=UartDev.USB):
    if not bmc.has_feature("uart_mux"):
        logger.warning("Board does not support UART mux, skipping...")
        return

    if uart not in [0, 1]:
        logger.error("Unsupported UART.")

    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix(f"UART{uart}_MUX_SEL0-O", socket_idx, 0), id.value & 0x1)
        bmc.gpio_set(gpio_prefix(f"UART{uart}_MUX_SEL1-O", socket_idx, 0), (id.value >> 1) & 0x1)
    logger.info(f"Configured UART{uart} mux to {id.name}")

# set JTAG die selection strap
def die_sel(cmd="die_sel_0"):
    try:
        die_num = int(cmd[8:], 0)
    except ValueError:
        logger.error("Malformed die_sel command")
        sys.exit(-1)

    if die_num > 0x7 or die_num < 0:
        logger.error("Die number must be a valid 3 bit mask (0-7)")
        sys.exit(-1)

    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("DIE_SEL0-B", socket_idx, 1), die_num & 0x1)
        bmc.gpio_set(gpio_prefix("DIE_SEL1-B", socket_idx, 1), (die_num >> 1) & 0x1)
        bmc.gpio_set(gpio_prefix("DIE_SEL2-B", socket_idx, 1), (die_num >> 2) & 0x1)
    logger.info(f"Applied die mask {die_num:x}")

# set usb mgmt0 mux configuration
def usb_mgmt_sel(id=UsbMgmt0Dev.MGMT0):
    if not bmc.has_feature("usb_mgmt_mux"):
        logger.warning("Board does not support USB MGMT0 mux, skipping...")
        return

    if id == UsbMgmt0Dev.J2002_CON:
        logger.warning("This action will disconnect MCU from BMC. Continue? (y/n)")
        response = input().strip().lower()
        if response != 'y':
            logger.error("Action cancelled, skipping...")
            return

    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("USB_MGMT0_MUX_SEL0-O", socket_idx, 0), id.value & 0x1)
        bmc.gpio_set(gpio_prefix("USB_MGMT0_MUX_SEL1-O", socket_idx, 0), (id.value >> 1) & 0x1)
    logger.info(f"Configured BMC USB MGMT0 mux to {id.name}")

# set usb mctp mux configuration (default to BMC)
def usb_mctp_sel(id=UsbMctpDev.BMC):
    if not bmc.has_feature("usb_mctp_mux"):
        logger.warning("Board does not support USB MCTP mux, skipping...")
        return

    for socket in socket_list:
        bmc.tca95xx_write(socket["mgmt0"], STRAPS_1, 0, 2, id.value & 0x1) # MCTP_USB2_MUX_SEL
    logger.info(f"Configured BMC USB MCTP mux to {id.name}")

# set usb2 mux configuration (default to BMC)
def usb2_sel(id=Usb2Dev.BMC):
    if not bmc.has_feature("usb2_mux"):
        if bmc.has_feature("usb_mctp_mux"):
            logger.warning("Board uses USB MCTP mux instead of USB2_P1 mux, remapping...")
            usb_mctp_sel(UsbMctpDev(id.value))
            return
        logger.warning("Board does not support USB2 mux, skipping...")
        return

    for socket_idx in range(len(socket_list)):
        bmc.gpio_set(gpio_prefix("USB2_P1_MUX_SEL-O", socket_idx, 0), id.value & 0x1)

    logger.info(f"Configured USB2 mux to {id.name}")

# set package id strap
def package_id_sel(id=PackageId.PRIMARY):
    global bmc

    if not bmc.has_feature("package_id"):
        logger.warning("Board does not support setting PACKAGE_ID strap, skipping...")
        return

    bmc.gpio_set(gpio_prefix("GP32_PACKAGE_R_ID0-B", 0, 1), id.value)

    logger.info(f"Setting PACKAGE_ID strap to {id.name}")

# set qspi mux configuration
def qspi_sel(qspi, target=QspiDev.FLASH):
    global bmc
    if not bmc.has_feature("qspi_mux"):
        logger.warning("Board does not support QSPI mux, skipping...")
        return

    if qspi not in [0, 1]:
        logger.error("Unsupported QSPI.")
        bmc.close()
        sys.exit(-1)

    try:
        # multisocket probably isn't relevant here as only eboard supported for now
        # but ya never know
        for socket_idx in range(len(socket_list)):
            if target == QspiDev.QSPI_BRIDGE:
                bmc.gpio_set(gpio_prefix("QSPI_LB_EN-O", socket_idx, 0), 0x1)
                bmc.gpio_set(gpio_prefix("QSPI_LB_MUX_SEL-O", socket_idx, 0), 0x1)
                bmc.gpio_set(gpio_prefix("QSPI0_MUX_EN-O", socket_idx, 0), 0x1)
                bmc.gpio_set(gpio_prefix("QSPI1_MUX_EN-0", socket_idx, 0), 0x1)
            elif target == QspiDev.FLASH:
                bmc.gpio_set(gpio_prefix("QSPI_LB_EN-O", socket_idx, 0), 0x0)
                bmc.gpio_set(gpio_prefix("QSPI_LB_MUX_SEL-O", socket_idx, 0), 0x0)
                if qspi == 0:
                    bmc.gpio_set(gpio_prefix("QSPI0_MUX_EN-O", socket_idx, 0), 0x0)
                    bmc.gpio_set("OOB_MUX_CTRL0-O", "0x1")
                else:
                    bmc.gpio_set(gpio_prefix("QSPI1_MUX_EN-O", socket_idx, 0), 0x0)
                    bmc.gpio_set(gpio_prefix("QSPI1_MUX1_SEL-O", socket_idx, 0), 0x0)
                    bmc.gpio_set("OOB_MUX_CTRL1-O", "0x0")
            elif target == QspiDev.OOBHUB:
                bmc.gpio_set(gpio_prefix("QSPI_LB_EN-O", socket_idx, 0), 0x0)
                bmc.gpio_set(gpio_prefix("QSPI_LB_MUX_SEL-O", socket_idx, 0), 0x0)
                if qspi == 0:
                    bmc.gpio_set(gpio_prefix("QSPI0_MUX_EN-O", socket_idx, 0), 0x0)
                    bmc.gpio_set("OOB_MUX_CTRL0-O", "0x0")
                else:
                    bmc.gpio_set(gpio_prefix("QSPI1_MUX_EN-O", socket_idx, 0), 0x0)
                    bmc.gpio_set(gpio_prefix("QSPI1_MUX1_SEL-O", socket_idx, 0), 0x0)
                    bmc.gpio_set("OOB_MUX_CTRL1-O", "0x1")
            elif target == QspiDev.SPI1:
                if qspi != 1:
                    logger.error("SPI1 is only supported on QSPI1")
                    bmc.close()
                    sys.exit(-1)
                bmc.gpio_set(gpio_prefix("QSPI_LB_EN-O", socket_idx, 0), 0x0)
                bmc.gpio_set(gpio_prefix("QSPI_LB_MUX_SEL-O", socket_idx, 0), 0x0)
                bmc.gpio_set(gpio_prefix("QSPI1_MUX_EN-O", socket_idx, 0), 0x0)
                bmc.gpio_set(gpio_prefix("QSPI1_MUX1_SEL-O", socket_idx, 0), 0x1)
    except:
        logger.error(f"Failed to set QSPI{qspi} to {target.name}--check MCU virtual I2C.")

    logger.info(f"Configured QSPI{qspi} to {target.name}")

# use default arguments to set to default
def set_defaults(skip_mux=False):
    logger.info("Resetting controls to defaults")
    boot_sel()
    rcm_sel()
    if not skip_mux:
        uart_sel(0)
        uart_sel(1)
        usb_mgmt_sel()
        usb2_sel() # does mctp for relevant plats
        qspi_sel(0)
        qspi_sel(1)
        package_id_sel()
    pull_rcm(1) # RCM off
    pull_ist(0) # Normal boot

# print gpio info
def gpio_info(concise=False, skip_mux=False):
    for idx, socket in enumerate(socket_list):
        if (bmc.has_feature("uart_mux") or bmc.has_feature("usb_mgmt_mux") or bmc.has_feature("usb2_mux") or
            bmc.has_feature("usb_mctp_mux") or bmc.has_feature("qspi_mux")) and not skip_mux:

            print(f"\n====== Socket {idx} MGMT 0 ======")

        if bmc.has_feature("uart_mux") and not skip_mux:
            uart_mux = bmc.tca95xx_read(socket["mgmt0"], UART_MUX_CTL, 0)
            uart0 = uart_mux & 0x3
            uart1 = (uart_mux >> 2) & 0x3

            if concise:
                print(f"UART0: {UartDev(uart0).name}")
                print(f"UART1: {UartDev(uart1).name}")
            else:
                print("" \
                    "UART0 -- mux -%s-------------%s- mux -- UART1\n" \
                    "           \\--%s- HDR   HDR -%s--/\n" \
                    "           \\--%s---- MCU ----%s--/\n" \
                    "           \\--%s---- USB ----%s--/\n" %
                    (("✅" if uart0 == UartDev.UART_BRIDGE.value else "❌"),
                     ("✅" if uart1 == UartDev.UART_BRIDGE.value else "❌"),
                     ("✅" if uart0 == UartDev.UART_HDR.value else "❌"),
                     ("✅" if uart1 == UartDev.UART_HDR.value else "❌"),
                     ("✅" if uart0 == UartDev.MCU.value else "❌"),
                     ("✅" if uart1 == UartDev.MCU.value else "❌"),
                     ("✅" if uart0 == UartDev.USB.value else "❌"),
                     ("✅" if uart1 == UartDev.USB.value else "❌"))
                )

        if bmc.has_feature("usb_mgmt_mux") and not skip_mux:
            usb_mux = bmc.tca95xx_read(socket["mgmt0"], USB_MUX_CTL, 0)
            usb0 = ((usb_mux >> 2) & 0x1) | ((usb_mux >> 6) & 0x2)

            if concise:
                print(f"USB0: {UsbMgmt0Dev(usb0).name}")
            else:
                print("" \
                    "USB_BMC_MGMT0 -- mux -%s-- BMC_MGMT0\n" \
                    "                   \\--%s-- BMC_MGMT1_HUB\n" \
                    "                   \\--%s-- J2002_CON\n" %
                    (("✅" if usb0 == UsbMgmt0Dev.MGMT0.value else "❌"),
                     ("✅" if usb0 == UsbMgmt0Dev.MGMT1_HUB.value else "❌"),
                     ("✅" if usb0 == UsbMgmt0Dev.J2002_CON.value else "❌"))
                )

        if bmc.has_feature("usb2_mux") and not skip_mux:
            usb_mux = bmc.tca95xx_read(socket["mgmt0"], USB_MUX_CTL, 0)
            usb2 = (usb_mux >> 6) & 0x1

            if concise:
                print(f"USB2: {Usb2Dev(usb2).name}")
            else:
                print("" \
                    "USB2_P1 -- mux -%s-- HOST (CONNECTOR)\n" \
                    "             \\--%s-- BMC\n" %
                    (("✅" if usb2 == Usb2Dev.HOST.value else "❌"),
                     ("✅" if usb2 == Usb2Dev.BMC.value else "❌"))
                )

        if bmc.has_feature("usb_mctp_mux") and not skip_mux:
            usb_mux = bmc.tca95xx_read(socket["mgmt0"], USB_MUX_CTL, 0)
            usb_mctp = (usb_mux >> 2) & 0x1
            if concise:
                print(f"USB3: {UsbMctpDev(usb_mctp).name}")
            else:
                print("" \
                    "USB3_HUB_MCTP_MGMT -- mux -%s-- HOST\n" \
                    "                        \\--%s-- BMC\n" %
                    (("✅" if usb_mctp == UsbMctpDev.HOST.value else "❌"),
                     ("✅" if usb_mctp == UsbMctpDev.BMC.value else "❌"))
                )

        if bmc.has_feature("qspi_mux") and not skip_mux:
            try:
                oob_mux = bmc.tca95xx_read(socket["mcu"], OOB_CTL, 1)
                oob_ctrl0 = (oob_mux >> 2) & 0x1
                oob_ctrl1 = (oob_mux >> 3) & 0x1
    
                qspi0_mux = bmc.tca95xx_read(socket["mgmt0"], UART_MUX_CTL, 0)
                qspi0_mux_en = (qspi0_mux >> 6) & 0x1
                qspi_lb_sel = (qspi0_mux >> 5) & 0x1
                qspi_lb_en = (qspi0_mux >> 4) & 0x1
    
                qspi1_mux = bmc.tca95xx_read(socket["mgmt0"], USB_MUX_CTL, 0)
                qspi1_mux_sel = (qspi1_mux >> 5) & 0x1
                qspi1_mux_en = (qspi1_mux >> 4) & 0x1
    
                if concise:
                    if ~oob_ctrl0 & ~qspi0_mux_en:
                        print("QSPI0: OOBHUB")
                    elif qspi_lb_en & qspi_lb_sel:
                        print("QSPI0: BRIDGE")
                    elif ~qspi_lb_en & ~qspi_lb_sel:
                        txt = "QSPI0:"
                        if oob_ctrl0 & ~qspi0_mux_en:
                            txt += " QSPI0_FLASH1"
                        txt += " QSPI0_FLASH2"
                        print(txt)
                    else:
                        print("QSPI0: UNKNOWN")
    
                    if ~qspi_lb_en & qspi_lb_sel:
                        print("QSPI1: BRIDGE")
                    elif oob_ctrl1 & ~qspi1_mux_en:
                        print("QSPI1: OOBHUB")
                    elif qspi_lb_en:
                        if ~qspi1_mux_en & ~qspi1_mux_sel:
                            print("QSPI1: QSPI1_FLASH")
                        elif ~qspi1_mux_en & qspi1_mux_sel:
                            print("QSPI1: SPI1")
                        else:
                            print("QSPI1: UNKNOWN")
                    else:
                        print("QSPI1: UNKNOWN")
                else:
                    print("" \
                        "         OOBHUB -%s-> mux -%s-> QSPI0_FLASH1\n" \
                        "                     %s\n" \
                        "QSPI0 -%s-> mux --%s--\n" \
                        "             %s       \\------> QSPI0_FLASH2\n" \
                        "             |\n" \
                        "             %s            /--%s-> SPI1 \n" \
                        "QSPI1 -%s-> mux -%s-%s-> mux -%s----%s-> mux -> QSPI1_FLASH\n" \
                        "                            OOBHUB -%s---/\n" %
                        (("✅" if oob_ctrl0 == 0 else "❌"),
                        ("✅" if qspi0_mux_en == 0 else "❌"),
                        ("✅" if oob_ctrl0 == 1 else "❌"),
                        ("✅" if qspi_lb_en == 0 else "❌"),
                        ("✅" if qspi_lb_sel == 0 else "❌"),
                        ("✅" if qspi_lb_sel == 1 else "❌"),
                        ("✅" if qspi_lb_sel == 1 else "❌"),
                        ("✅" if qspi1_mux_sel == 1 else "❌"),
                        ("✅" if qspi_lb_en == 0 else "❌"),
                        ("✅" if qspi_lb_sel == 0 else "❌"),
                        ("✅" if qspi1_mux_en == 0 else "❌"),
                        ("✅" if qspi1_mux_sel == 0 else "❌"),
                        ("✅" if oob_ctrl1 == 0 else "❌"),
                        ("✅" if oob_ctrl1 == 1 else "❌"))
                    )
            except:
                logger.error("QSPI mux read failed--check MCU virtual I2C")

        print(f"\n====== Socket {idx} MGMT 1 ======")

        # GPIOs
        mgmt1 = bmc.tca95xx_read(socket["mgmt1"], STRAPS_1, 0)

        if not concise:
            print("\n" \
                "           /------ CPU_IST_BOOT\n" \
                "          / /----- CPU_RECOVERY_TYPE[1:0]\n" \
                "         / /  /--- CPU_BOOT_DEV_SEL[2:0]\n" \
                "        / /  /  /- CPU_FORCED_RECOVERY_L\n" \
                "        |/\\ / \\/")
        print(f"Port 0: {mgmt1:08b}")

        mgmt1 = bmc.tca95xx_read(socket["mgmt1"], STRAPS_1, 1)

        if not concise:
            print("\n" \
                "           /----   NVJTAG_SEL_S\n" \
                "          / /----  CPU_BOOT_COMPLETE\n" \
                "         / / /---- DIE_SEL[2:0]\n" \
                "        / / /  /-- PACKAGE_R_ID[1:0]\n" \
                "        |/ /  / /- CPU_BOOT_CHAIN0\n" \
                "        ||/ \\/\\/")
        print(f"Port 1: {mgmt1:08b}")

if __name__ == "__main__":
    try:
        main()
    except KeyboardInterrupt:
        logger.warning("Interrupted, shutting down...")
        bmc.close()
        sys.exit(-1)

    sys.exit(0)
