#!/usr/bin/env -S uv run --script
#
# /// script
# requires-python = ">=3.14"
# dependencies = ["pyjson5", "numpy"]
# ///

from __future__ import annotations

import json
import logging
import struct
import sys
from argparse import ArgumentParser
from dataclasses import dataclass
from pathlib import Path
from typing import Sequence

import numpy as np
import pyjson5

REPO_ROOT = Path(__file__).parents[3]

TEXTURE_PAGE_WIDTH: int = 256
TEXTURE_PAGE_HEIGHT: int = 256
TEXTURE_PAGE_SIZE: int = TEXTURE_PAGE_WIDTH * TEXTURE_PAGE_HEIGHT

logger = logging.getLogger(__name__)

type JsonValue = (
    None
    | bool
    | int
    | float
    | str
    | list[JsonValue]
    | dict[str, JsonValue]
)


class _ByteReader:
    def __init__(self, data: bytes):
        self._data = data
        self._pos = 0

    def seek(self, pos: int) -> None:
        if pos < 0 or pos > len(self._data):
            raise ValueError("seek out of range")
        self._pos = pos

    def tell(self) -> int:
        return self._pos

    def read_u8(self) -> int:
        (value,) = struct.unpack_from("<B", self._data, self._pos)
        self._pos += 1
        return int(value)

    def read_s16(self) -> int:
        (value,) = struct.unpack_from("<h", self._data, self._pos)
        self._pos += 2
        return int(value)

    def read_u16(self) -> int:
        (value,) = struct.unpack_from("<H", self._data, self._pos)
        self._pos += 2
        return int(value)

    def read_s32(self) -> int:
        (value,) = struct.unpack_from("<i", self._data, self._pos)
        self._pos += 4
        return int(value)

    def read_u32(self) -> int:
        (value,) = struct.unpack_from("<I", self._data, self._pos)
        self._pos += 4
        return int(value)

    def skip(self, count: int) -> None:
        self.seek(self._pos + count)

    def slice(self, count: int) -> bytes:
        start = self._pos
        end = start + count
        self._pos = end
        return self._data[start:end]


def _argb1555_to_rgba8888(packed: int) -> tuple[int, int, int, int]:
    a1 = (packed >> 15) & 0x01
    r5 = (packed >> 10) & 0x1F
    g5 = (packed >> 5) & 0x1F
    b5 = packed & 0x1F
    a8 = 255 if a1 != 0 else 0
    r8 = (r5 << 3) | (r5 >> 2)
    g8 = (g5 << 3) | (g5 >> 2)
    b8 = (b5 << 3) | (b5 >> 2)
    return (r8, g8, b8, a8)


def _skip_tr3_rooms(r: _ByteReader) -> None:
    room_count = r.read_u16()
    for _ in range(room_count):
        r.skip(16)
        num = r.read_s32()
        r.skip(num * 2)
        num = r.read_u16()
        r.skip(num * 32)

        size_z = r.read_s16()
        size_x = r.read_s16()
        r.skip(int(size_z) * int(size_x) * 8)

        r.skip(4)
        num = r.read_u16()
        r.skip(num * 24)
        num = r.read_u16()
        r.skip(num * 20)
        r.skip(7)

    num = r.read_s32()
    r.skip(num * 2)


@dataclass(frozen=True, slots=True)
class TR3LevelItem:
    object_id: int
    x: int
    y: int
    z: int
    room_num: int


@dataclass(frozen=True, slots=True)
class TR3LevelData:
    palette_rgba: np.ndarray
    pages_rgba: list[np.ndarray]
    object_textures: list[dict[str, int | list[tuple[int, int]]]]
    anims: list[dict[str, int]]
    bones: list[dict[str, int | bool]]
    frame_data: np.ndarray
    mesh_blob: bytes
    mesh_offsets: list[int]
    objects: list[int]
    items: list[TR3LevelItem]


def read_tr3_level(level_path: str) -> TR3LevelData:
    with open(level_path, "rb") as f:
        level = f.read()

    r = _ByteReader(level)
    version = r.read_u32()
    if version != 0xFF080038 and version != 0xFF180038:
        raise ValueError(f"not a TR3 level file (version=0x{version:08X})")
    logger.debug("version ok @ %d", r.tell())

    pal_rgb = r.slice(256 * 3)
    r.skip(
        256 * 4
    )  # unused in this script; kept for correct offset progression
    palette_rgba = np.zeros((256, 4), dtype=np.uint8)
    palette_rgba[:, 0:3] = np.frombuffer(pal_rgb, dtype=np.uint8).reshape(
        256, 3
    )
    palette_rgba[:, 3] = 255
    logger.debug("palettes read @ %d", r.tell())

    num_pages = r.read_s32()
    if num_pages <= 0:
        raise ValueError("no texture pages")
    logger.debug("texture pages=%d @ %d", num_pages, r.tell())

    r.skip(num_pages * TEXTURE_PAGE_SIZE)
    pages_16 = r.slice(num_pages * TEXTURE_PAGE_SIZE * 2)
    logger.debug("pages read (16-bit) @ %d", r.tell())
    pages_rgba: list[np.ndarray] = []
    for page_idx in range(num_pages):
        start = page_idx * TEXTURE_PAGE_SIZE * 2
        pix = np.frombuffer(
            pages_16[start : start + TEXTURE_PAGE_SIZE * 2], dtype="<u2"
        )
        rgba = np.zeros(
            (TEXTURE_PAGE_HEIGHT, TEXTURE_PAGE_WIDTH, 4), dtype=np.uint8
        )
        rgba_flat = rgba.reshape(-1, 4)
        for i in range(TEXTURE_PAGE_SIZE):
            rgba_flat[i] = _argb1555_to_rgba8888(int(pix[i]))
        pages_rgba.append(rgba)

    r.skip(4)
    logger.debug("unused version skipped @ %d", r.tell())
    _skip_tr3_rooms(r)
    logger.debug("rooms skipped @ %d", r.tell())

    mesh_data_words = r.read_s32()
    logger.debug("mesh_data_words=%d @ %d", mesh_data_words, r.tell())
    mesh_blob = r.slice(mesh_data_words * 2)
    mesh_ptr_count = r.read_s32()
    logger.debug("mesh_ptr_count=%d @ %d", mesh_ptr_count, r.tell())
    mesh_offsets = [r.read_s32() for _ in range(mesh_ptr_count)]
    logger.debug("mesh_offsets read @ %d", r.tell())

    anim_count = r.read_s32()
    logger.debug("anim_count=%d @ %d", anim_count, r.tell())
    anims: list[dict[str, int]] = []
    for _ in range(anim_count):
        anims.append(
            {
                "frame_ofs": r.read_u32(),
                "interpolation": r.read_u8(),
                "frame_size": r.read_u8(),
                "current_anim_state": r.read_s16(),
                "velocity": r.read_s32(),
                "acceleration": r.read_s32(),
                "frame_base": r.read_s16(),
                "frame_end": r.read_s16(),
                "jump_anim_num": r.read_s16(),
                "jump_frame_num": r.read_s16(),
                "num_changes": r.read_s16(),
                "change_idx": r.read_s16(),
                "num_commands": r.read_s16(),
                "command_idx": r.read_s16(),
            }
        )
    logger.debug("anims read @ %d", r.tell())

    num = r.read_s32()
    logger.debug("anim changes count(words?)=%d @ %d", num, r.tell())
    r.skip(num * 6)
    num = r.read_s32()
    logger.debug("anim ranges count(words?)=%d @ %d", num, r.tell())
    r.skip(num * 8)
    num = r.read_s32()
    logger.debug("anim commands count(words?)=%d @ %d", num, r.tell())
    r.skip(num * 2)

    bone_total_int32 = r.read_s32()
    if bone_total_int32 < 0 or (bone_total_int32 % 4) != 0:
        raise ValueError("invalid anim bone block size")
    bone_count = bone_total_int32 // 4
    logger.debug("bone_count=%d @ %d", bone_count, r.tell())
    bones: list[dict[str, int | bool]] = []
    for _ in range(bone_count):
        flags = r.read_s32()
        bones.append(
            {
                "matrix_pop": (flags & 1) != 0,
                "matrix_push": (flags & 2) != 0,
                "pos_x": r.read_s32(),
                "pos_y": r.read_s32(),
                "pos_z": r.read_s32(),
            }
        )
    logger.debug("bones read @ %d", r.tell())

    frame_data_words = r.read_s32()
    logger.debug("frame_data_words=%d @ %d", frame_data_words, r.tell())
    frame_data = np.frombuffer(r.slice(frame_data_words * 2), dtype="<i2")
    logger.debug("frame_data read @ %d", r.tell())

    num_objects = r.read_s32()
    logger.debug("objects=%d @ %d", num_objects, r.tell())
    objects_by_game_id: dict[int, dict[str, int]] = {}
    for _ in range(num_objects):
        gid = r.read_s32()
        mesh_count = r.read_s16()
        mesh_idx = r.read_s16()
        bone_idx = r.read_s32() // 4
        frame_ofs = r.read_u32()
        anim_idx = r.read_s16()
        objects_by_game_id[gid] = {
            "mesh_count": int(mesh_count),
            "mesh_idx": int(mesh_idx),
            "bone_idx": int(bone_idx),
            "frame_ofs": int(frame_ofs),
            "anim_idx": int(anim_idx),
        }
    logger.debug("objects block read @ %d", r.tell())

    num_static = r.read_s32()
    logger.debug("static objects=%d @ %d", num_static, r.tell())
    r.skip(num_static * 32)
    logger.debug("static objects skipped @ %d", r.tell())

    num = r.read_s32()
    logger.debug("sprite textures=%d @ %d", num, r.tell())
    r.skip(num * 16)
    num = r.read_s32()
    logger.debug("sprite sequences=%d @ %d", num, r.tell())
    r.skip(num * 8)
    num = r.read_s32()
    logger.debug("cameras/sinks=%d @ %d", num, r.tell())
    r.skip(num * 16)
    num = r.read_s32()
    logger.debug("sound sources=%d @ %d", num, r.tell())
    r.skip(num * 16)

    box_count = r.read_s32()
    logger.debug("boxes=%d @ %d", box_count, r.tell())
    r.skip(box_count * 8)
    num = r.read_s32()
    logger.debug("overlaps=%d @ %d", num, r.tell())
    r.skip(num * 2)
    r.skip(box_count * 20)
    logger.debug("zones skipped @ %d", r.tell())

    num = r.read_s32()
    logger.debug("animated texture ranges=%d @ %d", num, r.tell())
    r.skip(num * 2)

    object_texture_count = r.read_s32()
    logger.debug("object textures=%d @ %d", object_texture_count, r.tell())
    object_textures: list[dict[str, int | list[tuple[int, int]]]] = []
    for _ in range(object_texture_count):
        draw_type = r.read_u16()
        tex_page = r.read_u16()
        uvs: list[tuple[int, int]] = []
        for __ in range(4):
            uvs.append((r.read_u16(), r.read_u16()))
        object_textures.append(
            {
                "draw_type": int(draw_type),
                "tex_page": int(tex_page),
                "uvs": uvs,
            }
        )
    logger.debug("object textures read @ %d", r.tell())

    num_items = r.read_s32()
    logger.debug("items=%d @ %d", num_items, r.tell())
    items: list[TR3LevelItem] = []
    for _ in range(num_items):
        object_id = r.read_s16()
        room_num = r.read_s16()
        x = r.read_s32()
        y = r.read_s32()
        z = r.read_s32()
        r.skip(2)  # y_rot
        r.skip(4)  # shade (value_1, value_2)
        r.skip(2)  # flags
        items.append(
            TR3LevelItem(
                object_id=int(object_id),
                x=int(x),
                y=int(y),
                z=int(z),
                room_num=int(room_num),
            )
        )
    logger.debug("items read @ %d", r.tell())

    return TR3LevelData(
        palette_rgba=palette_rgba,
        pages_rgba=pages_rgba,
        object_textures=object_textures,
        anims=anims,
        bones=bones,
        frame_data=frame_data,
        mesh_blob=mesh_blob,
        mesh_offsets=mesh_offsets,
        objects=sorted(objects_by_game_id.keys()),
        items=items,
    )


@dataclass(frozen=True, slots=True)
class Args:
    out: Path | None


def parse_args(argv: Sequence[str] | None = None) -> Args:
    parser = ArgumentParser()
    parser.add_argument(
        "-o",
        "--output",
        dest="out",
        type=Path,
        default=None,
        help="Write JSON output to this path (prints to stdout if omitted).",
    )
    ns = parser.parse_args(argv)
    return Args(out=ns.out)


def _json_dumps_limited_indent(
    value: JsonValue, *, indent: int, max_depth: int
) -> str:
    def encode(node: JsonValue, depth: int) -> str:
        if depth >= max_depth or not isinstance(node, (list, dict)):
            return json.dumps(
                node,
                separators=(",", ":"),
                ensure_ascii=False,
            )

        pad = " " * (indent * depth)
        child_pad = " " * (indent * (depth + 1))

        if isinstance(node, list):
            if not node:
                return "[]"
            inner = ",\n".join(
                f"{child_pad}{encode(item, depth + 1)}" for item in node
            )
            return f"[\n{inner}\n{pad}]"

        if not node:
            return "{}"
        inner = ",\n".join(
            f"{child_pad}{json.dumps(str(k), ensure_ascii=False)}: {encode(v, depth + 1)}"
            for k, v in node.items()
        )
        return f"{{\n{inner}\n{pad}}}"

    if max_depth < 0:
        raise ValueError("max_depth must be >= 0")
    if indent <= 0:
        raise ValueError("indent must be > 0")
    return encode(value, 0)


def main() -> None:
    logging.basicConfig(
        level=logging.INFO,
        format="%(levelname)s: %(message)s",
        stream=sys.stderr,
    )
    args: Args = parse_args()

    gf_path = REPO_ROOT / "data/trx/ship/games/tr3/gameflow.json5"
    strings_path = REPO_ROOT / "data/trx/ship/games/tr3/strings.json5"
    data = pyjson5.loads(gf_path.read_text())
    strings = pyjson5.loads(strings_path.read_text())

    levels: list[dict[str, object]] = []
    for idx, level in enumerate(data["levels"]):
        if level.get("type") == "gym":
            zone_num = -1
        else:
            zone_num = -1
            for zone_idx, entry in enumerate(data.get("globe_select_entries", [])):
                start = int(entry.get("start_level_ordinal", -1))
                end = int(entry.get("completion_level_ordinal", -1))
                if start < 0 or end < 0:
                    continue
                if start <= idx <= end:
                    zone_num = zone_idx
                    break

        level_path = REPO_ROOT / "test/trx/games/tr3/levels" / level["path"]
        logger.info("Reading %s", level_path)

        parsed = read_tr3_level(str(level_path))
        title = strings["levels"][idx]["title"]
        levels.append(
            {
                "title": title,
                "zone_num": zone_num,
                "objects": parsed.objects,
                "items": [
                    {
                        "object_id": item.object_id,
                        "x": item.x,
                        "y": item.y,
                        "z": item.z,
                        "room_num": item.room_num,
                    }
                    for item in parsed.items
                ],
            }
        )

    out_json = _json_dumps_limited_indent(levels, indent=4, max_depth=2) + "\n"
    if args.out is None:
        logger.info("No --out provided; writing JSON to stdout")
        print(out_json)
    else:
        args.out.parent.mkdir(parents=True, exist_ok=True)
        args.out.write_text(out_json, encoding="utf-8")
        print(args.out)


if __name__ == "__main__":
    main()
