#!/usr/bin/env python3
"""
Graphics decoder v1 for bit Generations: Orbital (GBA).

Reads the extracted asset image (assets/ regions concatenated in rom_off
order, sha1-checked against assets/manifest.json) and produces viewable
assets under assets/gfx/ using only the stdlib:

  - Palette candidates: runs of >=16 consecutive RGB555-plausible u16
    (bit15==0, not all zero) in the tail image -> palette_NN.pal + swatch.
  - Tile sheets: 4bpp 8x8 tiles (32B each), 16 tiles/row, rendered with a
    grayscale ramp plus recolored variants with the top palettes.

This is a v1 "viewable assets" pass, not an exact tilemap reconstruction:
sheet offsets are auto-picked by tile-likeness scoring (nibble entropy and
density), so sheets show real graphic data but not necessarily in-game
layout. Raw .bin blobs remain the build source of truth.

Usage:
  python3 tools/extract_assets.py            # first
  python3 tools/decode_gfx.py [--assets assets]
"""
import argparse
import json
import struct
import zlib
from collections import Counter
from math import log2
from pathlib import Path

ROOT = Path(__file__).resolve().parent.parent

TILE = 32          # bytes per 4bpp 8x8 tile
COLS = 16          # tiles per sheet row
SHEET_TILES = 256  # tiles per sheet (8KB)


def write_png(path: Path, w: int, h: int, rgb: bytes) -> None:
    def chunk(tag: bytes, data: bytes) -> bytes:
        c = struct.pack(">I", len(data)) + tag + data
        return c + struct.pack(">I", zlib.crc32(tag + data) & 0xFFFFFFFF)
    ihdr = struct.pack(">IIBBBBB", w, h, 8, 2, 0, 0, 0)
    raw = b"".join(b"\x00" + rgb[y * w * 3:(y + 1) * w * 3] for y in range(h))
    path.write_bytes(b"\x89PNG\r\n\x1a\n" + chunk(b"IHDR", ihdr)
                     + chunk(b"IDAT", zlib.compress(raw, 9))
                     + chunk(b"IEND", b""))


def rgb555(v: int) -> tuple:
    r, g, b = v & 31, (v >> 5) & 31, (v >> 10) & 31
    return (r * 255 // 31, g * 255 // 31, b * 255 // 31)


def load_image(assets: Path) -> bytes:
    man = json.loads((assets / "manifest.json").read_text())
    regs = [r for r in man["regions"] if not r.get("overlap")]
    regs.sort(key=lambda r: int(r["rom_off"], 16))
    blob = bytearray()
    cursor = 0
    for r in regs:
        off = int(r["rom_off"], 16)
        assert off == cursor, f"gap/overlap at {r['name']}: off {off:X} != {cursor:X}"
        data = (assets / r["name"]).read_bytes()
        assert len(data) == r["size"], f"size drift: {r['name']}"
        blob += data
        cursor += len(data)
    return bytes(blob)


def find_palettes(img: bytes, start: int, limit: int):
    """Runs of >=16 RGB555-plausible u16. Returns [(off, [u16...])]."""
    pals = []
    i = start
    n = len(img)
    while i < n - 1 and len(pals) < limit:
        v = struct.unpack_from("<H", img, i)[0]
        if v & 0x8000 or v == 0:
            i += 2
            continue
        j = i
        vals = []
        while j < n - 1 and len(vals) < 256:
            w = struct.unpack_from("<H", img, j)[0]
            if w & 0x8000:
                break
            vals.append(w)
            j += 2
        if len(vals) >= 16 and len(set(vals)) >= 32:
            pals.append((i, vals))
            i = j
        else:
            i += 2
    return pals


def tile_score(img: bytes, off: int) -> float:
    """Higher = more tile-like (dense nibbles, mid entropy)."""
    blk = img[off:off + 4096]
    if len(blk) < 4096:
        return -1.0
    zeros = blk.count(0)
    if zeros > 3000:
        return -1.0
    c = Counter(blk)
    ent = -sum(v / 4096 * log2(v / 4096) for v in c.values())
    if not 3.0 <= ent <= 6.5:
        return -1.0
    nib = Counter()
    for b in blk:
        nib[b & 15] += 1
        nib[b >> 4] += 1
    used = sum(1 for v in nib.values() if v > 64)
    return used + ent


def render_sheet(img: bytes, off: int, pal: list | None) -> tuple[int, int, bytes]:
    w, h = COLS * 8, (SHEET_TILES // COLS) * 8
    px = bytearray(w * h * 3)
    for t in range(SHEET_TILES):
        base = off + t * TILE
        if base + TILE > len(img):
            break
        for row in range(8):
            for col in range(8):
                b = img[base + row * 4 + col // 2]
                idx = (b >> 4) if col % 2 else (b & 15)
                if pal is not None and idx < len(pal):
                    r, g, bl = rgb555(pal[idx])
                else:
                    v = idx * 17
                    r = g = bl = v
                x = (t % COLS) * 8 + col
                y = (t // COLS) * 8 + row
                o = (y * w + x) * 3
                px[o:o + 3] = bytes((r, g, bl))
    return w, h, bytes(px)


def main() -> int:
    ap = argparse.ArgumentParser(description="Decode Orbital gfx v1")
    ap.add_argument("--assets", default=str(ROOT / "assets"))
    args = ap.parse_args()
    assets = Path(args.assets)
    img = load_image(assets)
    print(f"image: {len(img)} bytes")
    gfx = assets / "gfx"
    gfx.mkdir(parents=True, exist_ok=True)

    # Palettes from tail (skip header/code/rodata: 0x40000+)
    pals = find_palettes(img, 0x40000, 8)
    pal_info = []
    for k, (off, vals) in enumerate(pals):
        (gfx / f"palette_{k:02d}.pal").write_bytes(
            struct.pack(f"<{len(vals)}H", *vals))
        sw = bytearray()
        for v in vals:
            sw += bytes(rgb555(v))
        write_png(gfx / f"palette_{k:02d}.png", len(vals), 1, bytes(sw))
        pal_info.append({"id": k, "rom_off": f"0x{off:07X}",
                         "colors": len(vals)})
    print(f"palettes: {len(pals)}")

    # Tile-sheet offsets: fixed probe + top auto-scored windows
    scored = []
    for off in range(0x40000, len(img) - 8192, 0x2000):
        s = tile_score(img, off)
        if s > 0:
            scored.append((s, off))
    scored.sort(reverse=True)
    picked = [0x800000]
    for _, off in scored:
        if all(abs(off - p) >= 0x4000 for p in picked):
            picked.append(off)
        if len(picked) >= 5:
            break
    sheets = []
    for off in picked:
        if off + SHEET_TILES * TILE > len(img):
            continue
        w, h, px = render_sheet(img, off, None)
        name = f"sheet_{off:07X}.png"
        write_png(gfx / name, w, h, px)
        sheets.append({"rom_off": f"0x{off:07X}", "file": name,
                       "tiles": SHEET_TILES, "mode": "4bpp-gray"})
        for k, (_, vals) in enumerate(pals[:2]):
            w2, h2, px2 = render_sheet(img, off, vals)
            n2 = f"sheet_{off:07X}_pal{k:02d}.png"
            write_png(gfx / n2, w2, h2, px2)
            sheets.append({"rom_off": f"0x{off:07X}", "file": n2,
                           "tiles": SHEET_TILES,
                           "mode": f"4bpp-pal{k:02d}"})
    (gfx / "index.json").write_text(json.dumps(
        {"palettes": pal_info, "sheets": sheets}, indent=1) + "\n")
    print(f"sheets: {len(sheets)} -> {gfx}/index.json")
    return 0


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