#!/usr/bin/env python3
"""Reinstall firmware on a Dynamic Perception NMX over USB.

Guide: https://timelapsemaestro.com/nmx-firmware-recovery.html

Why not plain avrdude: the NMX's USB update loader (Dynamic Perception's
4 KB build of the LUFA CDC bootloader) mangles any address past 32 KB in its
"set address" command. Current avrdude sets the address before every
256-byte page, so most of the upper half of the firmware lands in the wrong
place and verification fails. This script sets the address ONCE, at 0, and
streams every page after it using the loader's own auto-increment, which
counts correctly. That is how Dynamic Perception's original updater got
away with it. The loader itself is write-protected and never touched.

Usage (NMX in update mode: hold E-Stop while plugging in USB):
    python3 nmx_flash.py nmx-0.77.hex      (Mac)
    py nmx_flash.py nmx-0.77.hex           (Windows)
The script finds the NMX's port by itself.
Options:
    --port PORT      use this port instead of searching (e.g. /dev/cu.usbmodem1101, COM4)
    --verify-only    compare the NMX against the file; write nothing
    --keep-settings  skip the factory settings reset

Mac and Linux need nothing extra. Windows needs pyserial:
    py -m pip install pyserial
"""

import argparse
import os
import sys
import time

PAGE = 256                   # SPM page size of the AT90USB1287 (loader block size)
BOOT_START = 0x1F000         # 4 KB loader at the top of flash; never written
SIGNATURE = b"\x82\x97\x1e"  # AT90USB1287, in the order the loader sends it
LOADER_ID = b"LUFACDC"
LOADER_USB_IDS = (0x03EB, 0x204A)  # USB vendor/product the loader reports (stock LUFA CDC)


class FlashError(Exception):
    pass


def load_hex(path):
    """Intel HEX -> {address: byte}. Handles record types 00, 01, 02 and 04."""
    mem, base = {}, 0
    try:
        f = open(path)
    except OSError as e:
        raise FlashError(f"Can't open {path}: {e.strerror}. Check the file name and folder.")
    with f:
        for n, line in enumerate(f, 1):
            line = line.strip()
            if not line:
                continue
            try:
                if not line.startswith(":"):
                    raise ValueError
                rec = bytes.fromhex(line[1:])
            except ValueError:
                raise FlashError(f"{path} line {n} is not Intel HEX. Is this the right file?")
            if sum(rec) & 0xFF:
                raise FlashError(f"{path} line {n} has a bad checksum. Try downloading it again.")
            addr, kind, data = (rec[1] << 8) | rec[2], rec[3], rec[4:4 + rec[0]]
            if kind == 0:
                for i, b in enumerate(data):
                    mem[base + addr + i] = b
            elif kind == 1:
                break
            elif kind == 2:
                base = ((data[0] << 8) | data[1]) << 4
            elif kind == 4:
                base = ((data[0] << 8) | data[1]) << 16
    if not mem:
        raise FlashError(f"{path} contains no firmware data.")
    if max(mem) >= BOOT_START:
        raise FlashError(f"{path} reaches into the loader area. It is not NMX firmware.")
    return mem


class PosixPort:
    """Mac/Linux serial port using only the standard library."""

    def __init__(self, path):
        import termios
        import tty
        self._termios = termios
        try:
            self.fd = os.open(path, os.O_RDWR | os.O_NOCTTY | os.O_NONBLOCK)
        except OSError as e:
            raise FlashError(f"Can't open {path}: {e.strerror}.")
        tty.setraw(self.fd)
        attrs = termios.tcgetattr(self.fd)
        attrs[2] |= termios.CLOCAL | termios.CREAD
        attrs[4] = attrs[5] = termios.B57600  # ignored over USB, set for tidiness
        termios.tcsetattr(self.fd, termios.TCSANOW, attrs)
        termios.tcflush(self.fd, termios.TCIOFLUSH)

    def write(self, data):
        import select
        view = memoryview(data)
        while view:
            select.select([], [self.fd], [], 5)
            try:
                view = view[os.write(self.fd, view):]
            except BlockingIOError:
                pass

    def read(self, n, timeout=3.0):
        import select
        out, deadline = bytearray(), time.monotonic() + timeout
        while len(out) < n:
            left = deadline - time.monotonic()
            if left <= 0 or not select.select([self.fd], [], [], left)[0]:
                break
            try:
                chunk = os.read(self.fd, n - len(out))
            except BlockingIOError:
                continue
            if not chunk:
                break
            out += chunk
        return bytes(out)

    def drain_input(self):
        time.sleep(0.2)
        self._termios.tcflush(self.fd, self._termios.TCIFLUSH)

    def close(self):
        os.close(self.fd)


class PySerialPort:
    """Windows (or anywhere pyserial is preferred)."""

    def __init__(self, path):
        try:
            import serial
        except ImportError:
            raise FlashError("This needs pyserial on Windows. Run:  py -m pip install pyserial")
        try:
            self.s = serial.Serial(path, 57600, timeout=3.0, write_timeout=10)
        except (serial.SerialException, OSError) as e:
            raise FlashError(f"Can't open {path}: {e}.")
        self.s.reset_input_buffer()

    def write(self, data):
        self.s.write(data)
        self.s.flush()

    def read(self, n, timeout=3.0):
        self.s.timeout = timeout
        return bytes(self.s.read(n))

    def drain_input(self):
        time.sleep(0.2)
        self.s.reset_input_buffer()

    def close(self):
        self.s.close()


def open_port(path):
    if os.name == "nt" or os.environ.get("NMX_FLASH_PYSERIAL"):
        return PySerialPort(path)
    return PosixPort(path)


def candidate_ports(windows=(os.name == "nt")):
    """Ports that could be the NMX's update loader."""
    try:
        from serial.tools import list_ports
    except ImportError:
        if windows:
            raise FlashError("This needs pyserial on Windows. Run:  py -m pip install pyserial")
    else:
        hits = sorted(p.device for p in list_ports.comports() if (p.vid, p.pid) == LOADER_USB_IDS)
        if hits or windows:
            return hits
    import glob
    return sorted(glob.glob("/dev/cu.usbmodem*") + glob.glob("/dev/ttyACM*"))


def find_port(wait=5.0):
    """The one port the NMX is on. Waits a few seconds in case it was just plugged in."""
    deadline = time.monotonic() + wait
    while True:
        ports = candidate_ports()
        if ports or time.monotonic() >= deadline:
            break
        time.sleep(0.5)
    if not ports:
        raise FlashError(
            "Couldn't find the NMX. Unplug it, then hold E-Stop while plugging USB back in "
            "(step 3), and run this again. If it still isn't found, try a different USB cable.")
    if len(ports) > 1:
        raise FlashError(
            f"Found more than one possible port: {', '.join(ports)}. Unplug other USB devices "
            f"and run this again, or choose one by adding --port and its name to the command.")
    return ports[0]


def ask(port, cmd, n, what):
    port.write(cmd)
    reply = port.read(n)
    if len(reply) != n:
        raise FlashError(f"The NMX stopped answering during {what}.")
    return reply


def expect_cr(port, cmd, what):
    if ask(port, cmd, 1, what) != b"\r":
        raise FlashError(f"The NMX refused {what}.")


def set_address_zero(port, what):
    # The ONLY address this script ever sends. Everything else uses auto-increment.
    expect_cr(port, b"A\x00\x00", what)


def progress(label, done, total):
    width = 40
    fill = width * done // total
    sys.stdout.write(f"\r  {label} [{'#' * fill}{'.' * (width - fill)}] {100 * done // total:3d}%")
    sys.stdout.flush()
    if done == total:
        sys.stdout.write("\n")


def connect(port):
    port.write(b"\x1b\x1b\x1b")  # sync bytes the loader ignores
    port.drain_input()
    port.write(b"S")
    if port.read(7, timeout=2.0) != LOADER_ID:
        raise FlashError(
            "The NMX isn't in update mode. Unplug it, then hold E-Stop while plugging "
            "USB back in (step 3), and run this again.")
    if ask(port, b"s", 3, "the chip check") != SIGNATURE:
        raise FlashError("Unexpected chip signature. This doesn't look like an NMX.")
    block = ask(port, b"b", 3, "the block-size check")
    if block[0:1] != b"Y" or (block[1] << 8 | block[2]) != PAGE:
        raise FlashError("The loader reports an unexpected block size. Stopping to be safe.")
    expect_cr(port, b"P", "entering update mode")


def page_count(mem):
    return (max(mem) // PAGE) + 1  # contiguous from address 0


def write_flash(port, mem):
    n = page_count(mem)
    set_address_zero(port, "the start of the firmware write")
    for p in range(n):
        page = bytes(mem.get(p * PAGE + i, 0xFF) for i in range(PAGE))
        expect_cr(port, b"B\x01\x00F" + page, f"writing page {p + 1} of {n}")
        progress("Writing  ", p + 1, n)


def verify_flash(port, mem):
    n = page_count(mem)
    set_address_zero(port, "the start of verification")
    bad = []
    for p in range(n):
        got = ask(port, b"g\x01\x00F", PAGE, f"reading page {p + 1} of {n}")
        for i, b in enumerate(got):
            a = p * PAGE + i
            if a in mem and mem[a] != b:
                bad.append(a)
        progress("Verifying", p + 1, n)
    return bad


def reset_settings(port):
    # Clears the "settings saved" flag (EEPROM byte 0) so the firmware
    # rebuilds factory defaults on its next start.
    set_address_zero(port, "the settings reset")
    expect_cr(port, b"B\x00\x01E\xff", "the settings reset")
    set_address_zero(port, "checking the settings reset")
    if ask(port, b"g\x00\x01E", 1, "checking the settings reset") != b"\xff":
        raise FlashError("The settings reset didn't stick.")


def main():
    ap = argparse.ArgumentParser(description="Reinstall Dynamic Perception NMX firmware over USB.")
    ap.add_argument("firmware", help="e.g. nmx-0.77.hex")
    ap.add_argument("--port", help="use this port instead of searching, e.g. /dev/cu.usbmodem1101 or COM4")
    ap.add_argument("--verify-only", action="store_true", help="compare the NMX against the file; write nothing")
    ap.add_argument("--keep-settings", action="store_true", help="skip the factory settings reset")
    args = ap.parse_args()

    port = None
    try:
        mem = load_hex(os.path.expanduser(args.firmware))
        print(f"Firmware: {os.path.basename(args.firmware)} ({len(mem)} bytes)")
        path = args.port or find_port()
        if not args.port:
            print(f"Found the NMX on {path}.")
        port = open_port(path)
        connect(port)
        print("Connected to the NMX update loader.")

        if not args.verify_only:
            write_flash(port, mem)
        bad = verify_flash(port, mem)
        if bad:
            print(f"\nVerify FAILED: {len(bad)} bytes differ (first at 0x{bad[0]:05X}).")
            if not args.verify_only:
                print("Please email this output to support@timelapsemaestro.com and we'll help.")
            return 1
        print(f"Verified: all {len(mem)} bytes match.")

        if args.verify_only:
            print("Verify-only run: nothing was written.")
        elif args.keep_settings:
            print("Settings left as they were (--keep-settings).")
        else:
            reset_settings(port)
            print("Saved settings cleared. The NMX will rebuild factory defaults on its next start.")

        port.write(b"L")
        port.read(1, timeout=1.0)
        port.write(b"E")  # leave the loader; the NMX restarts into its firmware
        port.read(1, timeout=1.0)
        if args.verify_only:
            print("\nDone. You can unplug the USB cable.")
        else:
            print("\nDone. Unplug the USB cable, reconnect your motors, and power the NMX up normally.")
        return 0
    except FlashError as e:
        print(f"\nError: {e}")
        return 1
    except KeyboardInterrupt:
        print("\nStopped. The update loader is untouched; just run the command again.")
        return 1
    finally:
        if port:
            port.close()


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