#!/usr/bin/env python3
"""
cromakga3d - KTX Compress (texture)
Converts crassets_raw/texture/**/*.png -> crassets/texture/**/*.ktx2 via the ktx CLI.
The CLI must be the same version as the libktx the engine links (see KTX_VERSION in
CMakeLists.txt / KtxVersion in sdk.props) — encoder and runtime are a matched pair.

Each source image requires a same-basename .crtex (TOML) contract; see
crassets_raw/texture/_template.crtex. Missing file / missing field / unknown field
is a hard error (fail-fast). Files whose basename starts with '_' are skipped.

[encode] fields are baked into the .ktx2 (self-describing). [runtime] fields are
injected as ktx2 key/value metadata (crtex_wrap / crtex_filter) so the .ktx2 is a
single self-contained artifact; the engine loader reads them back to set GL state.
"""

import os
import sys
import struct
import subprocess

if sys.version_info < (3, 11):
    sys.stderr.write("ktx_compress requires Python 3.11+ (tomllib).\n")
    sys.exit(1)
import tomllib

# -- Config -----------------------------------------------------------------
INPUT_ROOT = r"C:\workspace\cromakga3d\crassets_raw\texture"
OUT_ROOT   = r"C:\workspace\cromakga3d\crassets\texture"

IMAGE_EXTS = {".png"}

# Field -> allowed values. All fields mandatory; unknown fields rejected.
ENUM_FIELDS = {
    "compress":   {"uastc", "none"},
    "srgb":       {True, False},
    "mipmap":     {True, False},
    "mip_filter": {"box", "tent", "bell", "b-spline", "mitchell", "blackman",
                   "lanczos3", "lanczos4", "lanczos6", "lanczos12", "kaiser",
                   "gaussian", "catmullrom", "quadratic_interp",
                   "quadratic_approx", "quadratic_mix"},
}
RANGE_FIELDS = {
    "quality": (0, 4),
    "zstd":    (1, 22),
}
RUNTIME_ENUMS = {
    "wrap_s": {"clamp", "repeat", "mirror"},
    "wrap_t": {"clamp", "repeat", "mirror"},
    "filter": {"linear", "nearest"},
}
ENCODE_KEYS  = {"compress", "mipmap", "srgb", "quality", "mip_filter", "zstd"}
RUNTIME_KEYS = {"wrap_s", "wrap_t", "filter"}

# crtex wrap -> ktx --mipmap-wrap mode (GL names differ from ktx names)
MIPMAP_WRAP = {"clamp": "clamp", "repeat": "wrap", "mirror": "reflect"}

# -- ANSI colors ------------------------------------------------------------
os.system("")   # enable VT100 on Windows

CYAN, GREEN, RED, YELLOW, GRAY, WHITE, RESET = (
    "\033[96m", "\033[92m", "\033[91m", "\033[93m", "\033[90m", "\033[97m", "\033[0m")

def col(c, t):
    return c + t + RESET

def fmt_size(n):
    if n >= 1024 * 1024:
        s = "{:.1f} MB".format(n / (1024 * 1024))
    elif n >= 1024:
        s = "{:.1f} KB".format(n / 1024)
    else:
        s = str(n) + " B"
    return s.rjust(9)

def print_header():
    os.system("cls")
    print()
    print(col(CYAN, "  +==========================================+"))
    print(col(CYAN, "  |    cromakga3d  .  KTX Compress           |"))
    print(col(CYAN, "  +==========================================+"))
    print()
    print(col(GRAY, "  Input   :  " + INPUT_ROOT))
    print(col(GRAY, "  Output  :  " + OUT_ROOT))
    print()

def print_step(t):
    print(col(GRAY, "  ·  " + t))

def print_ok(t):
    print(col(GREEN, "  v  " + t))

def print_err(t):
    print(col(RED, "  x  " + t))

# -- Collection -------------------------------------------------------------
def collect(root, exts):
    result = []
    for dirpath, _dirs, files in os.walk(root):
        for name in files:
            if name.startswith("_"):
                continue
            if os.path.splitext(name)[1].lower() in exts:
                result.append(os.path.join(dirpath, name))
    result.sort()
    return result

def crtex_path(png_path):
    return os.path.splitext(png_path)[0] + ".crtex"

# -- Validation (fail-fast, batched) ----------------------------------------
def validate_contract(png_path, errors):
    """Parse + validate the .crtex; returns contract dict or None (errors appended)."""
    cpath = crtex_path(png_path)
    rel = os.path.relpath(png_path, INPUT_ROOT)
    if not os.path.isfile(cpath):
        errors.append("{}  -- missing .crtex".format(rel))
        return None
    try:
        with open(cpath, "rb") as f:
            doc = tomllib.load(f)
    except Exception as e:
        errors.append("{}  -- .crtex parse error: {}".format(rel, e))
        return None

    enc = doc.get("encode", {})
    run = doc.get("runtime", {})
    flat = {**enc, **run}

    missing = (ENCODE_KEYS | RUNTIME_KEYS) - flat.keys()
    unknown = flat.keys() - (ENCODE_KEYS | RUNTIME_KEYS)
    for m in sorted(missing):
        errors.append("{}  -- missing field: {}".format(rel, m))
    for u in sorted(unknown):
        errors.append("{}  -- unknown field: {}".format(rel, u))
    if missing or unknown:
        return None

    for k, allowed in ENUM_FIELDS.items():
        if flat[k] not in allowed:
            errors.append("{}  -- {} = {!r} not allowed".format(rel, k, flat[k]))
    for k, (lo, hi) in RANGE_FIELDS.items():
        v = flat[k]
        if not isinstance(v, int) or not (lo <= v <= hi):
            errors.append("{}  -- {} = {!r} out of [{},{}]".format(rel, k, v, lo, hi))
    for k, allowed in RUNTIME_ENUMS.items():
        if flat[k] not in allowed:
            errors.append("{}  -- {} = {!r} not allowed".format(rel, k, flat[k]))

    return flat

def find_orphans(errors):
    """A .crtex with no matching .png (leftover from a rename/delete) is an error."""
    for dirpath, _dirs, files in os.walk(INPUT_ROOT):
        for name in files:
            if name.startswith("_") or not name.endswith(".crtex"):
                continue
            base = os.path.join(dirpath, os.path.splitext(name)[0])
            if not any(os.path.isfile(base + e) for e in IMAGE_EXTS):
                rel = os.path.relpath(os.path.join(dirpath, name), INPUT_ROOT)
                errors.append("{}  -- orphan .crtex (no source image)".format(rel))

# -- ktx create -------------------------------------------------------------
def derive_mipmap_wrap(wrap_s, wrap_t):
    """--mipmap-wrap is a single mode; collapse per-axis wrap. Matching axes
    win; a clamp axis yields to the tiling axis (mip seam consistency); two
    differing tiling axes fall back to S. Runtime tiling stays per-axis via KV;
    this only affects generated mip edge texels."""
    s, t = MIPMAP_WRAP[wrap_s], MIPMAP_WRAP[wrap_t]
    if s == t:
        return s
    if t == "clamp":
        return s
    if s == "clamp":
        return t
    return s

def build_create_cmd(c, src, dst):
    fmt = "R8G8B8A8_SRGB" if c["srgb"] else "R8G8B8A8_UNORM"
    tf  = "srgb" if c["srgb"] else "linear"
    cmd = ["ktx", "create", "--format", fmt, "--assign-tf", tf]

    if c["compress"] == "uastc":
        cmd += ["--encode", "uastc", "--uastc-quality", str(c["quality"])]

    if c["mipmap"]:
        cmd += ["--generate-mipmap",
                "--mipmap-filter", c["mip_filter"],
                "--mipmap-wrap", derive_mipmap_wrap(c["wrap_s"], c["wrap_t"])]

    cmd += ["--zstd", str(c["zstd"]), src, dst]
    return cmd

# -- ktx2 key/value injection (see module docstring) ------------------------
KTX2_ID = bytes([0xAB, 0x4B, 0x54, 0x58, 0x20, 0x32,
                 0x30, 0xBB, 0x0D, 0x0A, 0x1A, 0x0A])

def _parse_kvd(block):
    block = bytes(block)   # slices must be hashable (dict keys) — bytearray isn't
    entries, i = [], 0
    while i < len(block):
        (n,) = struct.unpack_from("<I", block, i); i += 4
        kv = block[i:i + n]; i += n + ((-n) % 4)
        z = kv.find(b"\x00")
        entries.append((kv[:z], kv[z + 1:]))
    return entries

def _build_kvd(pairs):
    out = bytearray()
    for key, val in sorted(pairs, key=lambda e: e[0]):
        kv = key + b"\x00" + val
        out += struct.pack("<I", len(kv)) + kv + b"\x00" * ((-len(kv)) % 4)
    return bytes(out)

def inject_kv(path, kv):
    """Insert crtex_* pairs into the ktx2 KVD. Assumes zstd supercompression
    (no SGD, no mip padding) so all mip data shifts uniformly by delta."""
    data = bytearray(open(path, "rb").read())
    if data[:12] != KTX2_ID:
        raise ValueError("not a KTX2 file")
    superc     = struct.unpack_from("<I", data, 44)[0]
    levelCount = struct.unpack_from("<I", data, 40)[0]
    kvdOff     = struct.unpack_from("<I", data, 56)[0]
    kvdLen     = struct.unpack_from("<I", data, 60)[0]
    sgdOff     = struct.unpack_from("<Q", data, 64)[0]
    if superc != 2 or sgdOff != 0 or kvdOff == 0:
        raise ValueError("unexpected ktx2 layout (superc={}, sgd={}, kvd={})"
                         .format(superc, sgdOff, kvdOff))

    pairs = dict(_parse_kvd(data[kvdOff:kvdOff + kvdLen]))
    for k, v in kv.items():
        pairs[k.encode()] = v.encode() + b"\x00"
    new_kvd = _build_kvd(list(pairs.items()))
    delta = len(new_kvd) - kvdLen

    out = bytearray(data[:kvdOff]) + new_kvd + data[kvdOff + kvdLen:]
    struct.pack_into("<I", out, 60, len(new_kvd))          # kvdByteLength
    for i in range(max(levelCount, 1)):                    # level byteOffsets
        p = 80 + i * 24
        struct.pack_into("<Q", out, p, struct.unpack_from("<Q", out, p)[0] + delta)
    open(path, "wb").write(out)

def ktx_validate(path):
    r = subprocess.run(["ktx", "validate", path], capture_output=True, text=True)
    return r.returncode == 0, (r.stderr or r.stdout).strip()

# -- Main -------------------------------------------------------------------
def main():
    print_header()

    if not os.path.isdir(INPUT_ROOT):
        print_err("INPUT_ROOT not found: " + INPUT_ROOT)
        sys.exit(1)

    pngs = collect(INPUT_ROOT, IMAGE_EXTS)

    # Phase 1 - validate everything, report all errors at once.
    errors = []
    contracts = {}
    for png in pngs:
        c = validate_contract(png, errors)
        if c is not None:
            contracts[png] = c
    find_orphans(errors)

    if errors:
        print(col(RED, "  Validation failed ({} issue(s)):".format(len(errors))))
        print()
        for e in errors:
            print_err(e)
        print()
        sys.exit(1)

    if not pngs:
        print(col(YELLOW, "  no textures found."))
        print()
        return

    # Phase 2 - encode.
    ok = fail = 0
    for png in pngs:
        c = contracts[png]
        rel = os.path.relpath(png, INPUT_ROOT)
        dst = os.path.join(OUT_ROOT, os.path.splitext(rel)[0] + ".ktx2")
        os.makedirs(os.path.dirname(dst), exist_ok=True)

        print_step(rel + "  →  " + os.path.splitext(rel)[0] + ".ktx2")

        r = subprocess.run(build_create_cmd(c, png, dst),
                           capture_output=True, text=True)
        if r.returncode != 0:
            print_err("create failed -- " + (r.stderr.strip() or "(no output)"))
            fail += 1
            continue

        try:
            inject_kv(dst, {"crtex_wrap_s": c["wrap_s"],
                            "crtex_wrap_t": c["wrap_t"],
                            "crtex_filter": c["filter"]})
        except Exception as e:
            print_err("kv inject failed -- " + str(e))
            fail += 1
            continue

        valid, msg = ktx_validate(dst)
        if not valid:
            print_err("validate failed -- " + (msg or "(no output)"))
            fail += 1
            continue

        src_size, dst_size = os.path.getsize(png), os.path.getsize(dst)
        ratio = (1.0 - dst_size / src_size) * 100.0 if src_size > 0 else 0.0
        rc = GREEN if ratio >= 0 else YELLOW
        print(col(GREEN, "  v  done") +
              "   " + col(GRAY, fmt_size(src_size)) +
              "  →  " + col(WHITE, fmt_size(dst_size)) +
              "   " + col(rc, "{:+.1f}%".format(-ratio)))
        ok += 1

    print()
    print(col(GRAY, "  ------------------------------------------"))
    print(col(GREEN, "  v  OK    : " + str(ok)))
    if fail:
        print(col(RED, "  x  Fail  : " + str(fail)))
    print(col(GRAY, "  ------------------------------------------"))
    print()

    if fail:
        sys.exit(1)


if __name__ == "__main__":
    main()
