#!/usr/bin/env python3
"""Side-by-side debug: ROM bytes vs linked candidate bytes + disassembly.

Usage: python3 tools/dbg.py <c_file> <hexaddr> [flagset]
  flagset: interwork (default, shows both), nointerwork
"""
import json, re, subprocess, sys, tempfile
from pathlib import Path

ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT / "tools"))
from check_direct import check_function_direct, real_extent, load_functions, ROM, AGBCC, AS, LD, OBJCOPY, ROM_BASE

TB = str(ROOT / "arm-gnu-toolchain-13.3.rel1-darwin-arm64-arm-none-eabi" / "bin")
OBJDUMP = f"{TB}/arm-none-eabi-objdump"

def run(cmd, **kw):
    return subprocess.run(cmd, capture_output=True, text=True, **kw)

def _scratch_dir():
    d = ROOT / ".scratch" / "dbg"
    d.mkdir(parents=True, exist_ok=True)
    return str(d)

def disasm(blob: bytes, vma: int) -> list:
    with tempfile.NamedTemporaryFile(suffix=".bin", delete=False,
                                      dir=_scratch_dir()) as f:
        f.write(blob); p = f.name
    r = run([OBJDUMP, "-b", "binary", "-m", "arm", "-M", "force-thumb",
             f"--adjust-vma={vma:#x}", "-D", p])
    Path(p).unlink()
    lines = []
    for ln in r.stdout.splitlines():
        ln = ln.strip()
        if re.match(r"^[0-9a-f]+:\s+[0-9a-f ]+", ln):
            lines.append(ln)
    return lines

def build_bytes(c_path: str, addr: int, flags):
    from check_direct import load_functions
    import re as _re
    funcs = load_functions()
    names = {a: n for a, s, n in funcs}
    real_size, rom_bytes = real_extent(addr)
    with tempfile.TemporaryDirectory(dir=str(ROOT / ".scratch")) as td:
        td = Path(td)
        s_path, o_path, elf, bin_path = td/"f.s", td/"f.o", td/"f.elf", td/"f.bin"
        r = run([AGBCC, *flags, "-S", str(c_path), "-o", str(s_path)])
        if not s_path.exists():
            return None, f"agbcc fail: {r.stderr[:300]}"
        print("----- agbcc asm -----")
        print(s_path.read_text(errors="replace"))
        print("---------------------")
        r = run([AS, "-mcpu=arm7tdmi", str(s_path), "-o", str(o_path)])
        if r.returncode != 0:
            return None, f"as fail: {r.stderr[:300]}"
        s_text = s_path.read_text(errors="replace")
        ext_names = set(_re.findall(r"^\s*\.extern\s+(\w+)", s_text, _re.M))
        ext_names |= set(_re.findall(r"\bbl\s+([A-Za-z_]\w*)", s_text))
        ext_names |= set(_re.findall(r"^\s*\.word\s+([A-Za-z_]\w*)", s_text, _re.M))
        ext_names -= {Path(c_path).stem}
        ext_names = {e for e in ext_names if not e[0].isdigit()}
        known = {n: a for a, s, n in funcs}
        try:
            import json as _j
            for _r in _j.loads((ROOT / "build" / "stub_map.json").read_text()):
                if _r.get("alias"):
                    _a = _r["addr"] if isinstance(_r["addr"], int) else int(_r["addr"], 16)
                    known.setdefault(_r["name"], _a)
        except Exception:
            pass
        stub_lines = ["\t.thumb"]
        for en in sorted(ext_names):
            if not _re.match(r"^[A-Za-z_][A-Za-z0-9_]*$", en):
                continue
            stub_lines.append(f"\t.globl {en}")
            if en in known:
                stub_lines.append(f"\t.type {en}, %function")
                stub_lines.append("\t.thumb_func")
                stub_lines.append(f"\t.set {en}, 0x{(known[en] | 1):08X}")
            else:
                stub_lines.append(f"\t.set {en}, 0x08000000")
        stub_s = td / "stubs.s"
        stub_o = td / "stubs.o"
        stub_s.write_text("\n".join(stub_lines) + "\n")
        r = run([AS, "-mcpu=arm7tdmi", str(stub_s), "-o", str(stub_o)])
        if r.returncode != 0:
            return None, f"stub-as fail: {r.stderr[:300]}"
        ld_script = td / "link.ld"
        ld_script.write_text('OUTPUT_FORMAT("elf32-littlearm")\nOUTPUT_ARCH(arm)\nSECTIONS {\n'
            f"  .text 0x{addr:08X} : {{ *(.text) *(.text.*) }}\n  /DISCARD/ : {{ *(*) }}\n}}\n")
        r = run([LD, "-T", str(ld_script), str(o_path), str(stub_o), "-o", str(elf)])
        if r.returncode != 0:
            return None, f"ld fail: {r.stderr[:300]}"
        r = run([OBJCOPY, "-O", "binary", "-j", ".text", str(elf), str(bin_path)])
        if r.returncode != 0:
            return None, f"objcopy fail: {r.stderr[:300]}"
        comp = bin_path.read_bytes()
        try:
            nm = run([str(Path(AS).parent / "arm-none-eabi-nm"), str(elf)])
            for _ln in nm.stdout.splitlines():
                _m = _re.match(r"\s*([0-9a-fA-F]+)\s+\w\s+(\S+)", _ln)
                if _m and _m.group(2) == Path(c_path).stem:
                    _off = int(_m.group(1), 16) - addr
                    if _off >= 0:
                        comp = comp[_off:]
                    break
        except Exception:
            pass
        return comp, None

def main():
    c_path, addr = sys.argv[1], int(sys.argv[2], 16)
    real_size, rom_bytes = real_extent(addr)
    print(f"ROM {addr:#x} real_size={real_size}: {rom_bytes.hex()}")
    print("----- ROM disasm -----")
    for ln in disasm(rom_bytes, addr):
        print("R " + ln)
    # also show pool words after end (up to next func)
    funcs = load_functions()
    starts = sorted(a for a, _, _ in funcs)
    nxt = starts[starts.index(addr)+1] if addr in starts and starts.index(addr)+1 < len(starts) else addr+64
    rom = ROM.read_bytes()
    pool = rom[addr-ROM_BASE: nxt-ROM_BASE]
    print(f"----- raw through next func ({nxt:#x}, {len(pool)}B): {pool.hex()} -----")
    for flags in (("-O2","-mthumb-interwork"), ("-O2","-mno-thumb-interwork"), ("-O1","-mthumb-interwork"), ("-O1","-mno-thumb-interwork"), ("-O2","-mthumb-interwork","-fno-omit-frame-pointer"), ("-O2","-mno-thumb-interwork","-fno-omit-frame-pointer")):
        print(f"===== flags: {' '.join(flags)} =====")
        comp, err = build_bytes(c_path, addr, flags)
        if err:
            print(err); continue
        ct = comp[:real_size]
        print(f"got {ct.hex()}")
        print(f"rom {rom_bytes.hex()}")
        mark = "".join("^" if a!=b else " " for a,b in zip(ct, rom_bytes))
        print(f"    {mark}  ({sum(1 for a,b in zip(ct,rom_bytes) if a!=b)} diffs)")
        print("----- linked disasm -----")
        for ln in disasm(ct, addr):
            print("C " + ln)

if __name__ == "__main__":
    main()
