#!/usr/bin/env python3
"""Reproducible instruction-range comparator and ordered-layout builder.

This is the Phase 5 matching harness. It has four subcommands:

  range   compile one C candidate with the fingerprinted original toolchain,
          assemble it, extract its exact instruction range, and compare that
          range byte-for-byte against the validated original executable.
  plan    print the address-ordered build plan for a region registry.
  build   generate and link the ordered binary: header, data gaps taken from
          the original, and the registry's C regions compiled in address order.
  gate    build, then compare the whole rebuilt executable (cmp + SHA-1).

Cross-references to unmatched functions and globals are supplied as absolute
assembler definitions. `--defsym NAME=0xADDR` adds one on the command line;
`--symbols FILE` loads a tracked registry of `NAME<TAB>address` rows. They are
resolved by the **linker**, so the assembler emits `%hi`/`%lo` relocations and
the linker applies the HI16 carry adjustment the original toolchain used.

`cc1` output is passed through `maspsx` (ASPSX emulator) before GNU `as`, which
is what gives the original's non-reordered delay-slot and `la`/`addiu` forms.
The ASPSX version is pinned to the SDK banner value (`2.81`).

A region row may carry an optional fourth field of per-region flag overrides:
space-separated `key=value` tokens with key `cc1` or `as`, each value a
comma-separated flag list. The flags are appended to the effective toolchain
flags for that region only (so `cc1=-O0` overrides a global `-O2`).

The original executable and every generated artifact are caller-supplied or
written to a caller-selected fresh directory. This tool never embeds game bytes
in its own source or output and refuses to write into an existing directory.

Exit codes: 0 success/match, 1 mismatch, 2 usage or environment error.

Toolchain (identified in Phase 6, correcting Phase 5): PsyQ 4.0's `CC1PSX` reports
`GNU C 2.7.2.SN32.3.7.0002`. The open decompals/old-gcc `gcc-2.7.2-psx` build is
instruction-identical to it across 21 probe files. Flags: `-O2 -G0` (the macro
address form is the default; this compiler has no `-mno-split-addresses`).
Input to cc1 must be preprocessed. The SDK 4.0 assembler is ASPSX 2.56, which
maspsx emulates.
"""

from __future__ import annotations

import argparse
import hashlib
from dataclasses import dataclass, replace
from pathlib import Path
import re
import shutil
import struct
import subprocess
import sys
from typing import Sequence


REPO_ROOT = Path(__file__).resolve().parent.parent
DEFAULT_CC1 = REPO_ROOT / "tools/old-gcc/gcc-2.7.2-psx/cc1"
_BINUTILS = REPO_ROOT / "tools/mipsel-none-elf-binutils/prefix/usr/bin"
DEFAULT_AS = _BINUTILS / "mipsel-none-elf-as"
DEFAULT_LD = _BINUTILS / "mipsel-none-elf-ld"
DEFAULT_OBJCOPY = _BINUTILS / "mipsel-none-elf-objcopy"
DEFAULT_NM = _BINUTILS / "mipsel-none-elf-nm"

# A symbol whose name is an address is a placeholder for that address, so it
# needs no registry row: `func_80017AD4` resolves to 0x80017AD4. The convention
# is already used by the registry (`g_80122354`, `D_8012E2C8`). A wrong address
# cannot pass unnoticed -- the byte gate compares the whole binary -- so the only
# cost of an implicit symbol is a loud mismatch, never a silent one.
ADDRESS_SYMBOL = re.compile(r"^(?:func|D|g|lbl)_([0-9A-Fa-f]{8})$")

DEFAULT_CPP_FLAGS = ["-E", "-P", "-undef"]
DEFAULT_CC1_FLAGS = ["-quiet", "-O2", "-G0"]
DEFAULT_AS_FLAGS = ["-march=r3000", "-G0"]
DEFAULT_MASPSX = REPO_ROOT / "tools/maspsx/maspsx.py"
# The PsyQ 4.0 SDK banner reports `Psy-Q ASPSX version 2.56`.
DEFAULT_ASPSX_VERSION = "2.56"

EXE_MAGIC = b"PS-X EXE"
HEADER_SIZE = 0x800
PAYLOAD_LMA = 0x800


class ToolError(Exception):
    """A usage or environment problem; maps to exit code 2."""


# --------------------------------------------------------------------------
# small helpers
# --------------------------------------------------------------------------


def sha1_file(path: Path) -> str:
    digest = hashlib.sha1()
    with path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def run(command: Sequence[str], cwd: Path | None = None, stdout=None) -> None:
    completed = subprocess.run(
        list(command), cwd=cwd, stdout=stdout, stderr=subprocess.PIPE
    )
    if completed.returncode:
        text = completed.stderr.decode("utf-8", "replace").strip().splitlines()
        message = text[0] if text else "no diagnostic"
        raise ToolError(f"command failed ({completed.returncode}): {message}")


def run_filter(command: Sequence[str], source: Path, destination: Path) -> None:
    """Run a stdin-to-stdout filter (used for the maspsx stage)."""
    with source.open("rb") as stdin, destination.open("wb") as stdout:
        completed = subprocess.run(
            list(command), stdin=stdin, stdout=stdout, stderr=subprocess.PIPE
        )
    if completed.returncode:
        text = completed.stderr.decode("utf-8", "replace").strip().splitlines()
        message = text[0] if text else "no diagnostic"
        raise ToolError(f"command failed ({completed.returncode}): {message}")


def parse_address(text: str) -> int:
    try:
        return int(text, 0)
    except ValueError as exc:
        raise ToolError(f"not an address: {text!r}") from exc


def parse_hex_address(text: str, line_number: int) -> int:
    """A symbol-registry address: hexadecimal, with or without a 0x prefix."""
    try:
        address = int(text, 16)
    except ValueError as exc:
        raise ToolError(f"symbols line {line_number}: not a hex address: {text!r}") from exc
    if not 0 <= address <= 0xFFFFFFFF:
        raise ToolError(f"symbols line {line_number}: address out of range")
    return address


def require_file(path: Path, label: str) -> Path:
    if not path.is_file():
        raise ToolError(f"{label} is not a regular file: {path}")
    return path


def fresh_directory(path: Path, label: str) -> Path:
    if path.exists() or path.is_symlink():
        raise ToolError(f"{label} already exists: {path}")
    return path


# --------------------------------------------------------------------------
# pure logic (unit-tested without any external tool)
# --------------------------------------------------------------------------


@dataclass(frozen=True)
class PsxExe:
    entry: int
    text_address: int
    text_size: int

    @property
    def payload_end(self) -> int:
        return self.text_address + self.text_size


def parse_psx_exe(header: bytes) -> PsxExe:
    if len(header) < HEADER_SIZE:
        raise ToolError("executable is smaller than a PS-X EXE header")
    if header[:8] != EXE_MAGIC:
        raise ToolError("executable does not carry the PS-X EXE magic")
    entry, _gp, text_address, text_size = struct.unpack_from("<IIII", header, 0x10)
    if text_size == 0:
        raise ToolError("PS-X EXE header declares an empty payload")
    return PsxExe(entry, text_address, text_size)


@dataclass(frozen=True)
class Region:
    start: int
    end: int
    source: str
    cc1_flags: tuple[str, ...] = ()
    as_flags: tuple[str, ...] = ()
    no_gp: tuple[str, ...] = ()
    no_maspsx: bool = False
    maspsx_flags: tuple[str, ...] = ()


# A region's optional fourth field: space-separated `key=value` overrides.
# The value is a comma-separated flag list; the keys name the toolchain stage.
# `gp=-NAME` excludes a symbol from the registry's gp markers for this region
# only, because the access form is a property of the site, not of the symbol:
# 0x80121F84 is read gp-relative at 0x800A80BC and written absolutely at
# 0x8002D288 in this executable.
# `maspsx=off` runs the region without the ASPSX emulation stage, because
# maspsx's unconditional `nop` for a jump destroys the delay-slot fill that GNU
# `as` reorder mode performs on an expanded symbol store (0x80102B10, 0x800F8B6C,
# 0x800F3160), while the ASPSX `la`/`addiu` form needs maspsx (func_8002D2BC).
# The same key also takes Phase 10's opt-in maspsx modes, which are local additions
# to the pinned vendored tool (tools/patches/maspsx-phase10-r1r2.patch):
#   `maspsx=noreordernop` suppresses maspsx's unconditional reorder `nop` after a
#     branch/jump so GNU `as` can fill the slot itself (worker C's R1).
#   `maspsx=regread` additionally treats a following `jr`/`jalr` that uses the
#     loaded register as needing the load-delay `nop` (worker C's R2).
# Both default off, so every region without them compiles exactly as before.
_REGION_OPTION_KEYS = ("cc1", "as", "gp", "maspsx")
_MASPSX_MODES = {"off": "--off", "noreordernop": "--no-jump-slot-nop",
                  "regread": "--nop-on-reg-read"}


def parse_region_options(text: str, line_number: int) -> tuple[tuple[str, ...], tuple[str, ...], tuple[str, ...], bool, tuple[str, ...]]:
    """Parse the optional per-region override field.

    Returns `(cc1_flags, as_flags, no_gp, no_maspsx, maspsx_flags)`. `gp=-NAME` names a symbol
    the registry marks `gp` that this region accesses absolutely instead, and
    `maspsx=off` drops the ASPSX emulation stage for this region.
    """
    overrides: dict[str, tuple[str, ...]] = {}
    for token in text.split():
        if "=" not in token:
            raise ToolError(
                f"regions line {line_number}: expected 'key=value' override, got {token!r}"
            )
        key, _, value = token.partition("=")
        if key not in _REGION_OPTION_KEYS:
            raise ToolError(
                f"regions line {line_number}: unknown override key {key!r}"
            )
        if key in overrides:
            raise ToolError(
                f"regions line {line_number}: duplicate override key {key!r}"
            )
        flags = tuple(part for part in value.split(",") if part)
        if not flags:
            raise ToolError(f"regions line {line_number}: empty override for {key!r}")
        overrides[key] = flags
    no_gp: list[str] = []
    for name in overrides.get("gp", ()):
        if not name.startswith("-") or len(name) == 1:
            raise ToolError(
                f"regions line {line_number}: expected 'gp=-NAME', got {name!r}"
            )
        if not _SYMBOL_NAME.match(name[1:]):
            raise ToolError(
                f"regions line {line_number}: invalid symbol name {name[1:]!r}"
            )
        no_gp.append(name[1:])
    no_maspsx = False
    maspsx_flags: list[str] = []
    for value in overrides.get("maspsx", ()):
        if value not in _MASPSX_MODES:
            raise ToolError(
                f"regions line {line_number}: expected one of "
                f"{', '.join(sorted(_MASPSX_MODES))} for 'maspsx', got {value!r}"
            )
        if value == "off":
            no_maspsx = True
        else:
            maspsx_flags.append(_MASPSX_MODES[value])
    return (overrides.get("cc1", ()), overrides.get("as", ()), tuple(no_gp), no_maspsx,
            tuple(maspsx_flags))


def parse_regions(text: str) -> list[Region]:
    regions: list[Region] = []
    for number, raw in enumerate(text.splitlines(), 1):
        line = raw.split("#", 1)[0].strip()
        if not line:
            continue
        fields = line.split(None, 3)
        if len(fields) not in (3, 4):
            raise ToolError(
                f"regions line {number}: expected 'start end source [overrides]'"
            )
        start = parse_address(fields[0])
        end = parse_address(fields[1])
        if not 0 <= start < end:
            raise ToolError(f"regions line {number}: invalid range")
        cc1_flags, as_flags, no_gp, no_maspsx, maspsx_flags = ((), (), (), False, ())
        if len(fields) == 4:
            cc1_flags, as_flags, no_gp, no_maspsx, maspsx_flags = parse_region_options(fields[3], number)
        regions.append(Region(start, end, fields[2], cc1_flags, as_flags, no_gp, no_maspsx, maspsx_flags))
    regions.sort(key=lambda region: region.start)
    for left, right in zip(regions, regions[1:]):
        if right.start < left.end:
            raise ToolError(f"regions overlap at 0x{right.start:08X}")
    return regions


# A tracked symbol registry: `NAME<TAB>address[<TAB>gp]` rows naming absolute
# addresses. The optional `gp` marker forces a gp-relative access.
_SYMBOL_NAME = re.compile(r"[A-Za-z_.][A-Za-z0-9_.]*\Z")


@dataclass(frozen=True)
class SymbolTable:
    defsyms: tuple[str, ...] = ()
    gp_names: frozenset[str] = frozenset()


def parse_symbols(text: str) -> SymbolTable:
    """Parse a symbol registry into linker `--defsym` args and gp markers.

    Columns are `NAME<TAB>address[<TAB>gp]`. A `gp` marker means the original
    accessed the symbol `gp`-relative, so the harness rewrites its macro
    accesses to explicit `%gp_rel` (GNU `as` will not do this for a symbol whose
    section it cannot see).
    """
    defsyms: list[str] = []
    gp_names: set[str] = set()
    seen: set[str] = set()
    for number, raw in enumerate(text.splitlines(), 1):
        line = raw.split("#", 1)[0].strip()
        if not line:
            continue
        fields = line.split()
        if len(fields) not in (2, 3):
            raise ToolError(f"symbols line {number}: expected 'NAME address [gp]'")
        name = fields[0]
        if not _SYMBOL_NAME.match(name):
            raise ToolError(f"symbols line {number}: invalid symbol name {name!r}")
        if name in seen:
            raise ToolError(f"symbols line {number}: duplicate symbol {name!r}")
        address = parse_hex_address(fields[1], number)
        if len(fields) == 3:
            if fields[2] != "gp":
                raise ToolError(f"symbols line {number}: unknown marker {fields[2]!r}")
            gp_names.add(name)
        seen.add(name)
        defsyms.append(f"{name}=0x{address:X}")
    return SymbolTable(tuple(defsyms), frozenset(gp_names))


def validate_regions(regions: Sequence[Region], exe: PsxExe) -> None:
    for region in regions:
        if region.start < exe.text_address or region.end > exe.payload_end:
            raise ToolError(
                f"region 0x{region.start:08X}..0x{region.end:08X} leaves the payload"
            )


def payload_gaps(exe: PsxExe, regions: Sequence[Region]) -> list[tuple[int, int]]:
    """The payload ranges not covered by a C region, in address order."""
    gaps: list[tuple[int, int]] = []
    cursor = exe.text_address
    for region in regions:
        if region.start > cursor:
            gaps.append((cursor, region.start))
        cursor = region.end
    if cursor < exe.payload_end:
        gaps.append((cursor, exe.payload_end))
    return gaps


@dataclass(frozen=True)
class Diff:
    expected_length: int
    actual_length: int
    differing_bytes: int
    first_difference: int

    @property
    def identical(self) -> bool:
        return self.differing_bytes == 0


def compare_bytes(expected: bytes, actual: bytes) -> Diff:
    shared = min(len(expected), len(actual))
    first = -1
    differing = 0
    for index in range(shared):
        if expected[index] != actual[index]:
            differing += 1
            if first < 0:
                first = index
    differing += abs(len(expected) - len(actual))
    return Diff(len(expected), len(actual), differing, first)


@dataclass(frozen=True)
class LayoutItem:
    kind: str  # "asm" (data gap) or "c" (region source)
    start: int
    end: int
    source: Path  # gap .s path, or the region's C source
    object_name: str


def plan_layout(
    exe: PsxExe, regions: Sequence[Region], out: Path
) -> list[LayoutItem]:
    items: list[LayoutItem] = []
    for index, (start, end) in enumerate(payload_gaps(exe, regions)):
        items.append(
            LayoutItem("asm", start, end, out / f"gap_{index:03d}.s", f"gap_{index:03d}.o")
        )
    for index, region in enumerate(regions):
        items.append(
            LayoutItem(
                "c",
                region.start,
                region.end,
                Path(region.source),
                f"region_{index:03d}.o",
            )
        )
    items.sort(key=lambda item: item.start)
    return items


def linker_script(items: Sequence[LayoutItem], text_address: int) -> str:
    lines = [
        'OUTPUT_FORMAT("elf32-littlemips")',
        "ENTRY(_start)",
        "SECTIONS",
        "{",
        "  .header 0x0 : { *(.header) }",
        f"  .main 0x{text_address:X} : AT(0x{PAYLOAD_LMA:X})",
        "  {",
    ]
    for item in items:
        if item.kind == "asm":
            lines.append(f"    {item.object_name}(.data)")
        else:
            lines.append(f"    {item.object_name}(.text .rodata .data)")
    lines += [
        "  }",
        "  /DISCARD/ : { *(.MIPS.abiflags) *(.reginfo) *(.pdr) *(.gnu.attributes)"
        " *(.comment) *(.note*) *(.mdebug*) }",
        "}",
        "",
    ]
    return "\n".join(lines)


def gap_source(exe_path: Path, start: int, end: int) -> str:
    offset = PAYLOAD_LMA + (start - 0x80010000)
    return (
        '.section .data,"a",@progbits\n'
        f'.incbin "{exe_path}",{offset},{end - start}\n'
    )


def header_source(exe_path: Path) -> str:
    return (
        '.section .header,"a",@progbits\n'
        f'.incbin "{exe_path}",0,{HEADER_SIZE}\n'
    )


# --------------------------------------------------------------------------
# toolchain steps
# --------------------------------------------------------------------------


@dataclass(frozen=True)
class Toolchain:
    cpp: Path
    cc1: Path
    assembler: Path
    linker: Path
    objcopy: Path
    cpp_flags: Sequence[str]
    cc1_flags: Sequence[str]
    as_flags: Sequence[str]
    defsyms: Sequence[str]
    gp_symbols: frozenset[str] = frozenset()
    maspsx: Path | None = None
    maspsx_flags: tuple[str, ...] = ()
    aspsx_version: str = DEFAULT_ASPSX_VERSION
    nm: Path | None = None


# A symbol macro access that can be forced gp-relative.
_GP_ACCESS = re.compile(
    r"^(?P<indent>\s*)(?P<op>lw|lh|lhu|lb|lbu|lwl|lwr|sw|sh|sb|swl|swr)"
    r"\s+(?P<rt>\$[A-Za-z0-9]+),\s*(?P<sym>[A-Za-z_.][A-Za-z0-9_.]*)"
    r"(?P<addend>[+-][0-9]+)?\s*$"
)
_GP_LA = re.compile(
    r"^(?P<indent>\s*)la\s+(?P<rt>\$[A-Za-z0-9]+),\s*"
    r"(?P<sym>[A-Za-z_.][A-Za-z0-9_.]*)(?P<addend>[+-][0-9]+)?\s*$"
)


def rewrite_gp_accesses(text: str, gp_symbols: frozenset[str]) -> str:
    """Force `%gp_rel` access for symbols the original read through `gp`."""
    if not gp_symbols:
        return text
    lines: list[str] = []
    for line in text.splitlines():
        match = _GP_ACCESS.match(line)
        if match and match.group("sym") in gp_symbols:
            symbol = match.group("sym") + (match.group("addend") or "")
            lines.append(
                f"{match.group('indent')}{match.group('op')} "
                f"{match.group('rt')},%gp_rel({symbol})($gp)"
            )
            continue
        match = _GP_LA.match(line)
        if match and match.group("sym") in gp_symbols:
            symbol = match.group("sym") + (match.group("addend") or "")
            lines.append(
                f"{match.group('indent')}addiu {match.group('rt')},"
                f"$gp,%gp_rel({symbol})"
            )
            continue
        lines.append(line)
    return "\n".join(lines) + ("\n" if text.endswith("\n") else "")


def compile_c(source: Path, out_object: Path, work: Path, tools: Toolchain) -> None:
    """Preprocess, compile, ASPSX-emulate and assemble one C source.

    Symbols are left undefined so the assembler emits `%hi`/`%lo` relocations;
    the linker resolves them (see `link_object_bytes` and `_build`).
    """
    require_file(source, "C source")
    work.mkdir(parents=True, exist_ok=True)
    preprocessed = work / (out_object.stem + ".i")
    assembly = work / (out_object.stem + ".s")
    with preprocessed.open("wb") as handle:
        run([str(tools.cpp), *tools.cpp_flags, str(source)], stdout=handle)
    run([str(tools.cc1), *tools.cc1_flags, str(preprocessed), "-o", str(assembly)])
    if tools.maspsx is not None:
        transformed = work / (out_object.stem + ".maspsx.s")
        run_filter(
            [sys.executable, str(tools.maspsx),
             f"--aspsx-version={tools.aspsx_version}",
             *tools.maspsx_flags],
            assembly, transformed,
        )
        assembly = transformed
    if tools.gp_symbols:
        rewritten = work / (out_object.stem + ".gp.s")
        rewritten.write_text(
            rewrite_gp_accesses(assembly.read_text(encoding="utf-8"), tools.gp_symbols),
            encoding="utf-8",
        )
        assembly = rewritten
    run([str(tools.assembler), *tools.as_flags, "-o", str(out_object), str(assembly)])


def assemble_asm(source: Path, out_object: Path, tools: Toolchain) -> None:
    run([str(tools.assembler), *tools.as_flags, "-o", str(out_object), str(source)])


def single_linker_script(start: int) -> str:
    """A minimal script placing one object's `.text` at `start`.

    Using an explicit script (rather than `-Ttext`) avoids the default linker
    script's own symbol definitions, which would otherwise shadow a user symbol
    such as `__bss_start`.
    """
    return "\n".join([
        'OUTPUT_FORMAT("elf32-littlemips")',
        "SECTIONS",
        "{",
        f"  .text 0x{start:X} : {{ *(.text) }}",
        "  /DISCARD/ : { *(.MIPS.abiflags) *(.reginfo) *(.pdr) *(.gnu.attributes)"
        " *(.comment) *(.note*) *(.mdebug*) }",
        "}",
        "",
    ])


def link_object_bytes(
    object_path: Path, start: int, work: Path, tools: Toolchain
) -> bytes:
    """Link one object at `start` and return its relocated `.text` bytes."""
    script = work / "single.ld"
    script.write_text(single_linker_script(start), encoding="ascii")
    elf = work / "single.elf"
    binary = work / "single.bin"
    command = [str(tools.linker), "-T", str(script), "--no-check-sections"]
    for symbol in tools.defsyms:
        command += ["--defsym", symbol]
    command += ["-o", str(elf), str(object_path)]
    run(command)
    run([str(tools.objcopy), "-O", "binary", "--only-section=.text",
         str(elf), str(binary)])
    return binary.read_bytes()


def extract_text_bytes(object_path: Path, destination: Path, tools: Toolchain) -> bytes:
    run([str(tools.objcopy), "-O", "binary", "--only-section=.text",
         str(object_path), str(destination)])
    return destination.read_bytes()


# A symbol that never exists: `--keep-global-symbol` then localizes everything else.
LOCALIZE_SENTINEL = "__sf3_keep_no_global_symbol"


def localize_symbols(object_path: Path, tools: Toolchain) -> None:
    """Make a region object's symbols local.

    Region objects are placed by the linker script, not by symbol name, and any
    cross-reference is supplied as an absolute `--defsym`. Localizing therefore
    changes nothing about the emitted bytes, and it lets one shared source file
    be instantiated for several regions without a duplicate-symbol clash.
    """
    run([str(tools.objcopy), f"--keep-global-symbol={LOCALIZE_SENTINEL}",
         str(object_path)])


# --------------------------------------------------------------------------
# subcommands
# --------------------------------------------------------------------------


def command_range(args: argparse.Namespace) -> int:
    exe_path = require_file(args.exe, "executable")
    source = require_file(args.source, "C source")
    if args.end <= args.start:
        raise ToolError("--end must be greater than --start")
    work = fresh_directory(args.work, "work directory")

    data = exe_path.read_bytes()
    exe = parse_psx_exe(data)
    if args.start < exe.text_address or args.end > exe.payload_end:
        raise ToolError("requested range leaves the payload")

    work.mkdir(parents=True)
    object_path = work / "candidate.o"
    tools = args.toolchain
    compile_c(source, object_path, work, tools)
    tools = replace(tools, defsyms=resolve_undefined_symbols(
        tools.defsyms, undefined_symbols(object_path, tools), str(source)))
    candidate = link_object_bytes(object_path, args.start, work, tools)

    expected_length = args.end - args.start
    if len(candidate) != expected_length:
        print(f"range=0x{args.start:08X}..0x{args.end:08X}")
        print(f"expected_bytes={expected_length}")
        print(f"candidate_bytes={len(candidate)}")
        print("result=LENGTH-MISMATCH")
        return 1

    offset = PAYLOAD_LMA + (args.start - exe.text_address)
    expected = data[offset:offset + expected_length]
    diff = compare_bytes(expected, candidate)

    print(f"range=0x{args.start:08X}..0x{args.end:08X}")
    print(f"candidate_bytes={len(candidate)}")
    print(f"differing_bytes={diff.differing_bytes}")
    if diff.identical:
        print("result=MATCH")
        return 0
    print(f"first_difference=0x{args.start + diff.first_difference:08X}")
    print("result=DIFF")
    return 1


def _read_regions(path: Path) -> list[Region]:
    require_file(path, "region registry")
    return parse_regions(path.read_text(encoding="utf-8"))


def command_plan(args: argparse.Namespace) -> int:
    exe_path = require_file(args.exe, "executable")
    exe = parse_psx_exe(exe_path.read_bytes())
    regions = _read_regions(args.regions)
    validate_regions(regions, exe)
    out = args.out if args.out else Path(".")
    for item in plan_layout(exe, regions, out):
        print(f"{item.kind}\t0x{item.start:08X}\t0x{item.end:08X}\t{item.source}\t{item.object_name}")
    return 0


def _build(args: argparse.Namespace) -> tuple[Path, Path]:
    exe_path = require_file(args.exe, "executable")
    exe = parse_psx_exe(exe_path.read_bytes())
    regions = _read_regions(args.regions)
    validate_regions(regions, exe)
    for region in regions:
        require_file(Path(region.source), f"C source for 0x{region.start:08X}")

    out = fresh_directory(args.out, "output directory")
    out.mkdir(parents=True)
    work = out / "work"
    work.mkdir()

    items = plan_layout(exe, regions, out)
    resolved_exe = exe_path.resolve()
    header = out / "header.s"
    header.write_text(header_source(resolved_exe), encoding="ascii")
    for item in items:
        if item.kind == "asm":
            item.source.write_text(
                gap_source(resolved_exe, item.start, item.end), encoding="ascii"
            )

    tools = args.toolchain
    # Symbols an address-named placeholder can satisfy are added as the objects
    # that reference them are built; anything else fails loudly, naming the
    # symbol, instead of leaving a bare linker error.
    link_defsyms = list(tools.defsyms)
    by_start = {region.start: region for region in regions}
    objects: list[str] = ["header.o"]
    assemble_asm(header, out / "header.o", tools)
    for item in items:
        object_path = out / item.object_name
        if item.kind == "asm":
            assemble_asm(item.source, object_path, tools)
        else:
            region = by_start[item.start]
            region_tools = replace(
                tools,
                cc1_flags=[*tools.cc1_flags, *region.cc1_flags],
                as_flags=[*tools.as_flags, *region.as_flags],
                gp_symbols=tools.gp_symbols - frozenset(region.no_gp),
                maspsx=None if region.no_maspsx else tools.maspsx,
                maspsx_flags=region.maspsx_flags,
            )
            compile_c(item.source, object_path, work, region_tools)
            localize_symbols(object_path, tools)
            link_defsyms = resolve_undefined_symbols(
                link_defsyms, undefined_symbols(object_path, tools), item.source)
        objects.append(item.object_name)

    script = out / "link.ld"
    script.write_text(linker_script(items, exe.text_address), encoding="ascii")
    elf = out / "scus_946_40.elf"
    link_command = [str(tools.linker), "-T", script.name, "--no-check-sections"]
    for symbol in link_defsyms:
        link_command += ["--defsym", symbol]
    link_command += ["-o", elf.name, *objects]
    run(link_command, cwd=out)
    rebuilt = out / "scus_946_40.rebuilt"
    run([str(tools.objcopy), "-O", "binary", elf.name, rebuilt.name], cwd=out)
    return exe_path, rebuilt


def region_summary(regions: Sequence[Region]) -> list[str]:
    """Report how much C the build actually contains."""
    lines = [f"c_regions={len(regions)}"]
    if not regions:
        lines.append("c_regions_note=none: build contains no C (data baseline only)")
    return lines


def command_build(args: argparse.Namespace) -> int:
    exe_path, rebuilt = _build(args)
    regions = _read_regions(args.regions)
    for line in region_summary(regions):
        print(line)
    print(f"rebuilt={rebuilt}")
    print(f"rebuilt_bytes={rebuilt.stat().st_size}")
    print(f"rebuilt_sha1={sha1_file(rebuilt)}")
    print(f"original_sha1={sha1_file(exe_path)}")
    return 0


def command_gate(args: argparse.Namespace) -> int:
    exe_path, rebuilt = _build(args)
    regions = _read_regions(args.regions)
    expected = exe_path.read_bytes()
    actual = rebuilt.read_bytes()
    diff = compare_bytes(expected, actual)
    for line in region_summary(regions):
        print(line)
    print(f"rebuilt={rebuilt}")
    print(f"rebuilt_bytes={len(actual)}")
    print(f"original_bytes={len(expected)}")
    print(f"differing_bytes={diff.differing_bytes}")
    print(f"rebuilt_sha1={sha1_file(rebuilt)}")
    print(f"original_sha1={sha1_file(exe_path)}")
    if not diff.identical:
        print("result=DIFF")
        return 1
    if args.expect_sha1 and sha1_file(rebuilt) != args.expect_sha1.lower():
        print("result=SHA1-MISMATCH")
        return 1
    print("result=MATCH")
    return 0


# --------------------------------------------------------------------------
# argument parsing
# --------------------------------------------------------------------------


def add_toolchain_arguments(parser: argparse.ArgumentParser) -> None:
    parser.add_argument("--cpp", type=Path, default=None,
                        help="preprocessor (default: the one on PATH)")
    parser.add_argument("--cc1", type=Path, default=DEFAULT_CC1)
    parser.add_argument("--assembler", type=Path, default=DEFAULT_AS)
    parser.add_argument("--linker", type=Path, default=DEFAULT_LD)
    parser.add_argument("--objcopy", type=Path, default=DEFAULT_OBJCOPY)
    parser.add_argument("--nm", type=Path, default=DEFAULT_NM,
                        help="nm, used to resolve address-named symbols implicitly")
    parser.add_argument("--maspsx", type=Path, default=DEFAULT_MASPSX,
                        help="ASPSX emulator run between cc1 and the assembler")
    parser.add_argument("--no-maspsx", action="store_true",
                        help="assemble cc1 output directly, without maspsx")
    parser.add_argument("--no-jump-slot-nop", action="store_true",
                        help="maspsx mode: suppress the unconditional reorder nop after a "
                             "branch/jump so GNU as can fill the slot (region: maspsx=noreordernop)")
    parser.add_argument("--nop-on-reg-read", action="store_true",
                        help="maspsx mode: also treat a following jr/jalr that uses the loaded "
                             "register as needing the load-delay nop (region: maspsx=regread)")
    parser.add_argument("--aspsx-version", default=DEFAULT_ASPSX_VERSION,
                        help="ASPSX version for maspsx (default: the SDK's 2.81)")
    parser.add_argument("--cpp-flag", action="append", default=[],
                        help="extra preprocessor flag (repeatable)")
    parser.add_argument("--cc1-flag", action="append", default=[],
                        help="extra cc1 flag, replacing the defaults when given")
    parser.add_argument("--as-flag", action="append", default=[],
                        help="extra assembler flag, replacing the defaults when given")
    parser.add_argument("--defsym", action="append", default=[],
                        help="assembler --defsym NAME=0xADDR (repeatable)")
    parser.add_argument("--no-gp", action="append", default=[], dest="no_gp",
                        help="do not treat NAME as gp-relative despite its registry marker (repeatable)")
    parser.add_argument("--symbols", type=Path, default=None,
                        help="tracked symbol registry file (NAME<TAB>address rows)")


def undefined_symbols(object_path: Path, tools: Toolchain) -> set[str]:
    """The symbols one assembled object still needs defined.

    Read from the object itself rather than guessed from its C source, so a
    symbol the source defines but never references is not mistaken for one that
    needs resolving.
    """
    if tools.nm is None:
        return set()
    completed = subprocess.run(
        [str(tools.nm), "-u", str(object_path)],
        stdout=subprocess.PIPE, stderr=subprocess.PIPE,
    )
    if completed.returncode:
        text = completed.stderr.decode("utf-8", "replace").strip().splitlines()
        message = text[0] if text else "no diagnostic"
        raise ToolError(f"nm failed ({completed.returncode}): {message}")
    names: set[str] = set()
    for line in completed.stdout.decode("utf-8", "replace").splitlines():
        fields = line.split()
        if fields:
            names.add(fields[-1])
    return names


def resolve_undefined_symbols(defsyms: Sequence[str], undefined: set[str],
                              label: str) -> list[str]:
    """Add an address-named placeholder for each undefined symbol; fail loudly otherwise.

    A name that is neither in the registry nor shaped like an address cannot be
    resolved, and the link would fail with a bare `ld` message. Failing here
    instead names the symbol and says what to do about it.
    """
    known = {defsym.split("=", 1)[0] for defsym in defsyms}
    resolved = list(defsyms)
    unresolved: list[str] = []
    for name in sorted(undefined):
        if name in known:
            continue
        match = ADDRESS_SYMBOL.match(name)
        if match is not None:
            resolved.append(f"{name}=0x{int(match.group(1), 16):X}")
            known.add(name)
        else:
            unresolved.append(name)
    if unresolved:
        listed = ", ".join(unresolved)
        raise ToolError(
            f"unresolved symbol(s) referenced by {label}: {listed}; add a row to the symbol "
            "registry (NAME<TAB>address[<TAB>gp]) or name the symbol func_XXXXXXXX / "
            "D_XXXXXXXX so it resolves to that address"
        )
    return resolved


def resolve_toolchain(args: argparse.Namespace) -> Toolchain:
    cpp = args.cpp
    if cpp is None:
        located = shutil.which("cpp")
        if not located:
            raise ToolError("no preprocessor found; pass --cpp")
        cpp = Path(located)
    defsyms = list(args.defsym)
    gp_names: set[str] = set()
    if args.symbols is not None:
        symbols_path = require_file(args.symbols, "symbol registry")
        table = parse_symbols(symbols_path.read_text(encoding="utf-8"))
        defsyms = list(table.defsyms) + defsyms
        gp_names |= table.gp_names
    maspsx = None
    if not args.no_maspsx:
        maspsx = require_file(args.maspsx, "maspsx")
    return Toolchain(
        cpp=require_file(cpp, "preprocessor"),
        cc1=require_file(args.cc1, "cc1"),
        assembler=require_file(args.assembler, "assembler"),
        linker=require_file(args.linker, "linker"),
        objcopy=require_file(args.objcopy, "objcopy"),
        nm=require_file(args.nm, "nm") if args.nm is not None else None,
        cpp_flags=args.cpp_flag or DEFAULT_CPP_FLAGS,
        cc1_flags=args.cc1_flag or DEFAULT_CC1_FLAGS,
        as_flags=args.as_flag or DEFAULT_AS_FLAGS,
        defsyms=defsyms,
        gp_symbols=frozenset(gp_names - set(getattr(args, "no_gp", ()))),
        maspsx=maspsx,
        maspsx_flags=tuple(
            flag for flag, enabled in (
                ("--no-jump-slot-nop", getattr(args, "no_jump_slot_nop", False)),
                ("--nop-on-reg-read", getattr(args, "nop_on_reg_read", False))) if enabled),
        aspsx_version=args.aspsx_version,
    )


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__,
                                     formatter_class=argparse.RawDescriptionHelpFormatter)
    subparsers = parser.add_subparsers(dest="command", required=True)

    range_parser = subparsers.add_parser("range", help="compare one C candidate's range")
    range_parser.add_argument("--exe", required=True, type=Path)
    range_parser.add_argument("--source", required=True, type=Path)
    range_parser.add_argument("--start", required=True, type=parse_address)
    range_parser.add_argument("--end", required=True, type=parse_address)
    range_parser.add_argument("--work", required=True, type=Path)
    add_toolchain_arguments(range_parser)
    range_parser.set_defaults(handler=command_range)

    plan_parser = subparsers.add_parser("plan", help="print the ordered build plan")
    plan_parser.add_argument("--exe", required=True, type=Path)
    plan_parser.add_argument("--regions", required=True, type=Path)
    plan_parser.add_argument("--out", type=Path, default=None)
    plan_parser.set_defaults(handler=command_plan)

    for name, handler, help_text in (
        ("build", command_build, "build the ordered binary"),
        ("gate", command_gate, "build, then compare the whole executable"),
    ):
        sub = subparsers.add_parser(name, help=help_text)
        sub.add_argument("--exe", required=True, type=Path)
        sub.add_argument("--regions", required=True, type=Path)
        sub.add_argument("--out", required=True, type=Path)
        add_toolchain_arguments(sub)
        if name == "gate":
            sub.add_argument("--expect-sha1", default=None)
        sub.set_defaults(handler=handler)

    return parser


def main(argv: Sequence[str] | None = None) -> int:
    parser = build_parser()
    args = parser.parse_args(argv)
    try:
        if hasattr(args, "cc1"):
            args.toolchain = resolve_toolchain(args)
        return args.handler(args)
    except ToolError as exc:
        print(f"error: {exc}", file=sys.stderr)
        return 2
    except OSError as exc:
        print(f"error: {exc}", file=sys.stderr)
        return 2


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