# -*- coding: utf-8 -*-

import vxi11
import socket
import sys
import time
import re
import threading
from collections import OrderedDict
import pprint
import os

HOST = "169.254.58.10"
DEVICE = "gpib0,12"

REGIONS = ["IPZ", "SECN", "PRIM"]

# Safety/timeouts
SOCKET_TIMEOUT_S = 6.0
HARD_QUERY_TIMEOUT_S = 12.0
RETRIES_PER_QUERY = 3
RECONNECT_EVERY_N_QUERIES = 96
INTER_QUERY_DELAY_S = 0.02
AFTER_RST_DELAY_S = 0.5

# Discovery caps (offset units used by DUMP? region,<offset>)
DISCOVERY_MAX_SPAN = {
    "IPZ": 16384,
    "SECN": 65536,
    "PRIM": 32768,
}

STEP_CANDIDATES = (16, 8)
LOW_BYTES_PER_DUMP_EXPECTED = 8


def _to_bytes(x) -> bytes:
    if isinstance(x, bytes):
        return x
    if isinstance(x, str):
        return x.encode("ascii", errors="replace")
    return str(x).encode("ascii", errors="replace")


def safe_close(inst):
    try:
        inst.close()
    except Exception:
        pass


def connect():
    socket.setdefaulttimeout(SOCKET_TIMEOUT_S)
    inst = vxi11.Instrument(HOST, DEVICE)
    try:
        inst.timeout = SOCKET_TIMEOUT_S
    except Exception:
        pass
    return inst


def write(inst, cmd: str):
    if not cmd.endswith("\n"):
        cmd += "\n"
    inst.write(cmd)


def _query_worker(inst, cmd: str, read_len: int, out: dict):
    try:
        write(inst, cmd)
        out["resp"] = _to_bytes(inst.read(read_len))
        out["err"] = None
    except Exception as e:
        out["resp"] = b""
        out["err"] = e


def hard_query(inst, cmd: str, read_len: int = 512, hard_timeout_s: float = HARD_QUERY_TIMEOUT_S):
    out = {"resp": None, "err": None}
    t = threading.Thread(target=_query_worker, args=(inst, cmd, read_len, out), daemon=True)
    t.start()
    t.join(hard_timeout_s)
    if t.is_alive():
        return None, TimeoutError("Hard timeout on {!r}".format(cmd))
    if out["err"] is not None:
        return None, out["err"]
    if not out["resp"]:
        return None, TimeoutError("Empty response on {!r}".format(cmd))
    return out["resp"], None


def query_retry(inst, cmd: str, read_len: int = 512):
    last_err = None
    for _attempt in range(1, RETRIES_PER_QUERY + 1):
        resp, err = hard_query(inst, cmd, read_len=read_len)
        if err is None and resp is not None:
            return inst, resp
        last_err = err
        safe_close(inst)
        time.sleep(0.15)
        inst = connect()
    raise RuntimeError("Query failed after {} tries: {!r} (last error: {})".format(
        RETRIES_PER_QUERY, cmd, last_err
    ))


def parse_dump_words(line: bytes) -> list[str]:
    s = line.decode("ascii", errors="ignore")
    return re.findall(r"\b[0-9A-Fa-f]{4}\b", s)


def get_lower_bytes_from_dump(line: bytes) -> bytes:
    words = parse_dump_words(line)
    if len(words) < 9:
        raise ValueError("Unexpected DUMP? response (need >=9 hex groups): {!r}".format(line))

    out = bytearray()
    for w in words[1:9]:
        out.append(int(w[2:4], 16))
    return bytes(out)


def expected_out_bytes(span: int, step: int, bytes_per_dump: int) -> int:
    return (span // step) * bytes_per_dump


def detect_step(inst, region: str) -> int:
    inst, r0 = query_retry(inst, "DUMP? {},0".format(region), 512)
    w0 = parse_dump_words(r0)
    if not w0 or w0[0].upper() != "0000":
        return 16

    for step in STEP_CANDIDATES:
        inst, r1 = query_retry(inst, "DUMP? {},{}".format(region, step), 512)
        w1 = parse_dump_words(r1)
        if w1 and w1[0].upper() == "{:04X}".format(step):
            return step

    return 16


def try_offset(inst, region: str, offset: int):
    try:
        inst, resp = query_retry(inst, "DUMP? {},{}".format(region, offset), 512)
        chunk = get_lower_bytes_from_dump(resp)
        if len(chunk) != LOW_BYTES_PER_DUMP_EXPECTED:
            return inst, False
        return inst, True
    except Exception:
        return inst, False


def discover_span(inst, region: str, step: int, max_span: int) -> int:
    print("\nDiscovering {}: step={}, max_span={}".format(region, step, max_span))

    inst, ok0 = try_offset(inst, region, 0)
    if not ok0:
        print("{}: not readable at offset 0 (skipping)".format(region))
        return 0

    # Exponential search
    last_good = 0
    hi = step
    while hi <= max_span:
        inst, ok = try_offset(inst, region, hi)
        print("{} probe offset={} -> {}".format(region, hi, "OK" if ok else "FAIL"))
        if ok:
            last_good = hi
            hi *= 2
            time.sleep(INTER_QUERY_DELAY_S)
            continue
        break

    if hi > max_span:
        print("{}: no failure up to max_span; using span={}".format(region, max_span))
        return max_span

    first_bad = hi
    lo = last_good

    # Binary search
    lo_n = lo // step
    hi_n = first_bad // step

    while hi_n - lo_n > 1:
        mid_n = (lo_n + hi_n) // 2
        mid = mid_n * step
        inst, ok = try_offset(inst, region, mid)
        print("{} bisect offset={} -> {}".format(region, mid, "OK" if ok else "FAIL"))
        if ok:
            lo_n = mid_n
        else:
            hi_n = mid_n
        time.sleep(INTER_QUERY_DELAY_S)

    span = hi_n * step
    print("{}: discovered span={} (first bad offset)".format(region, span))
    return span


def dump_region(inst, region: str, span: int, step: int, filename: str, log_filename: str):
    if span == 0:
        return inst, 0

    inst, resp0 = query_retry(inst, "DUMP? {},0".format(region), 512)
    bytes_per_dump = len(get_lower_bytes_from_dump(resp0))

    exp = expected_out_bytes(span, step, bytes_per_dump)
    print("\nDumping {} -> {!r} [span={}, step={}, expected_out={} bytes]".format(
        region, filename, span, step, exp
    ))

    written = 0
    qcount = 0

    with open(filename, "wb") as f, open(log_filename, "w", encoding="utf-8") as log:
        for offset in range(0, span, step):
            if qcount > 0 and (qcount % RECONNECT_EVERY_N_QUERIES) == 0:
                safe_close(inst)
                inst = connect()
                time.sleep(0.1)

            log.write("{},{}\n".format(region, offset))
            log.flush()

            inst, resp = query_retry(inst, "DUMP? {},{}".format(region, offset), 512)
            chunk = get_lower_bytes_from_dump(resp)
            f.write(chunk)
            written += len(chunk)
            qcount += 1

            pct = int((offset * 100) / span) if span else 100
            sys.stdout.write("\r{}% complete   (offset={})".format(pct, offset))
            sys.stdout.flush()
            time.sleep(INTER_QUERY_DELAY_S)

    sys.stdout.write("\r{} done. Wrote {} bytes.                    \n".format(region, written))
    sys.stdout.flush()
    return inst, written


def concat_files(output_filename: str, inputs: list[str]):
    with open(output_filename, "wb") as fout:
        for fn in inputs:
            with open(fn, "rb") as fin:
                fout.write(fin.read())


def main():
    pp = pprint.PrettyPrinter(indent=2)
    inst = connect()

    try:
        write(inst, "*RST")
        time.sleep(AFTER_RST_DELAY_S)

        inst, idn_raw = query_retry(inst, "*IDN?", 255)
        idn = idn_raw.decode("ascii", errors="replace").strip()
        parts = [x.strip() for x in idn.split(",")]
        if len(parts) < 4:
            raise ValueError("Unexpected *IDN? response: {!r}".format(idn))

        id_, model, serial, firmware = parts[0], parts[1], parts[2], parts[3]
        pp.pprint(OrderedDict(id=id_, model=model, serial=serial, firmware=firmware))

        for r in REGIONS:
            try:
                inst, p = query_retry(inst, "DUMP? {},0".format(r), 512)
                print("Probe DUMP? {},0 -> {}".format(r, p[:80]))
            except Exception as e:
                print("Probe DUMP? {},0 failed: {}".format(r, e))

        discovered = {}
        for r in REGIONS:
            try:
                step = detect_step(inst, r)
                span = discover_span(inst, r, step=step, max_span=DISCOVERY_MAX_SPAN.get(r, 16384))
                discovered[r] = (step, span)
            except Exception as e:
                print("{}: discovery failed: {}".format(r, e))
                discovered[r] = (16, 0)

        print("\nDiscovered:")
        pp.pprint(OrderedDict((r, {"step": discovered[r][0], "span": discovered[r][1]}) for r in REGIONS))

        region_files = []
        bytes_written = OrderedDict()

        for r in REGIONS:
            step, span = discovered[r]
            if span == 0:
                print("Skipping {} (span=0)".format(r))
                continue
            fn = "{}.{}.nvram".format(serial, r.lower())
            logfn = "{}.{}.log.txt".format(serial, r)
            inst, w = dump_region(inst, r, span, step, fn, logfn)
            region_files.append(fn)
            bytes_written[r] = w

        out = "{}.nvram.bin".format(serial)
        concat_files(out, region_files)

        print("\nFinished! Created {!r}".format(out))
        for r, w in bytes_written.items():
            print("{}: {} bytes".format(r, w))
        print("Output file size: {} bytes".format(os.path.getsize(out)))
        print("Included files: {}".format(region_files))

        return 0

    finally:
        safe_close(inst)


if __name__ == "__main__":
    raise SystemExit(main())
