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

from __future__ import annotations

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

import pyjson5


REPO_ROOT = Path(__file__).parents[3]
TR4_VERSION = 0x00345254

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_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]


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


@dataclass(frozen=True, slots=True)
class TR4LevelData:
    objects: list[int]
    items: list[TR4LevelItem]


@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 _read_chunk(r: _ByteReader, name: str) -> bytes:
    uncompressed_size = r.read_u32()
    compressed_size = r.read_u32()
    payload = r.slice(compressed_size)
    try:
        out = zlib.decompress(payload)
    except zlib.error as exc:
        raise ValueError(f"failed to inflate TR4 chunk {name}") from exc
    if len(out) != uncompressed_size:
        raise ValueError(
            f"TR4 chunk {name} inflated to {len(out)} bytes, "
            f"expected {uncompressed_size}"
        )
    return out


def _skip_tr4_room_mesh(r: _ByteReader) -> None:
    mesh_length_words = r.read_u32()
    r.skip(mesh_length_words * 2)


def _skip_tr4_rooms(r: _ByteReader) -> None:
    room_count = r.read_s16()
    logger.debug("rooms=%d @ %d", room_count, r.tell())
    for _ in range(room_count):
        r.skip(16)
        _skip_tr4_room_mesh(r)

        portal_count = r.read_s16()
        r.skip(portal_count * 32)

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

        r.skip(4)
        light_count = r.read_s16()
        r.skip(light_count * 46)

        static_mesh_count = r.read_s16()
        r.skip(static_mesh_count * 20)

        r.skip(2)  # flipped room
        r.skip(2)  # flags
        r.skip(3)  # water scheme, reverb info, alternate group

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


def _skip_tr4_object_meshes(r: _ByteReader) -> None:
    mesh_data_words = r.read_s32()
    r.skip(mesh_data_words * 2)
    mesh_ptr_count = r.read_s32()
    r.skip(mesh_ptr_count * 4)


def _skip_tr4_anims(r: _ByteReader) -> None:
    anim_count = r.read_s32()
    r.skip(anim_count * 40)


def _skip_counted_records(
    r: _ByteReader, *, record_size: int, count_size: int = 4
) -> int:
    count = r.read_s32() if count_size == 4 else r.read_s16()
    r.skip(count * record_size)
    return count


def _skip_tr4_anim_data(r: _ByteReader) -> None:
    _skip_counted_records(r, record_size=6)
    _skip_counted_records(r, record_size=8)
    command_count = r.read_s32()
    r.skip(command_count * 2)

    bone_block_size = r.read_s32()
    r.skip(bone_block_size * 4)

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


def _read_tr4_objects(r: _ByteReader) -> list[int]:
    object_count = r.read_s32()
    logger.debug("objects=%d @ %d", object_count, r.tell())
    objects: list[int] = []
    for _ in range(object_count):
        object_id = r.read_s32()
        objects.append(object_id)
        r.skip(14)
    return sorted(objects)


def _skip_tr4_static_objects(r: _ByteReader) -> None:
    _skip_counted_records(r, record_size=32)


def _skip_tr4_sprite_textures(r: _ByteReader) -> None:
    signature = r.slice(3)
    if signature != b"SPR":
        raise ValueError(f"unexpected TR4 sprite texture signature: {signature!r}")
    _skip_counted_records(r, record_size=16)


def _skip_tr4_sprite_sequences(r: _ByteReader) -> None:
    _skip_counted_records(r, record_size=8)


def _skip_tr4_cameras_and_sinks(r: _ByteReader) -> None:
    _skip_counted_records(r, record_size=16)


def _skip_tr4_flyby_cameras(r: _ByteReader) -> None:
    count = r.read_s16()
    r.skip(2)
    r.skip(count * 40)


def _skip_tr4_sound_sources(r: _ByteReader) -> None:
    _skip_counted_records(r, record_size=16)


def _skip_tr4_pathing_data(r: _ByteReader) -> None:
    box_count = r.read_s32()
    r.skip(box_count * 8)
    overlap_count = r.read_s32()
    r.skip(overlap_count * 2)
    r.skip(box_count * 20)


def _skip_tr4_animated_texture_ranges(r: _ByteReader) -> None:
    data_size = r.read_s32()
    r.skip(data_size * 2)
    r.skip(1)


def _skip_tr4_object_textures(r: _ByteReader) -> None:
    signature = r.slice(3)
    if signature != b"TEX":
        raise ValueError(f"unexpected TR4 object texture signature: {signature!r}")
    _skip_counted_records(r, record_size=38)


def _read_tr4_items(r: _ByteReader) -> list[TR4LevelItem]:
    item_count = r.read_s32()
    logger.debug("items=%d @ %d", item_count, r.tell())
    items: list[TR4LevelItem] = []
    for _ in range(item_count):
        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(2)  # shade
        r.skip(2)  # ocb
        r.skip(2)  # flags
        items.append(
            TR4LevelItem(
                object_id=object_id,
                x=x,
                y=y,
                z=z,
                room_num=room_num,
            )
        )
    return items


def _read_tr4_ai_items(r: _ByteReader) -> list[TR4LevelItem]:
    item_count = r.read_s32()
    logger.debug("ai items=%d @ %d", item_count, r.tell())
    items: list[TR4LevelItem] = []
    for _ in range(item_count):
            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)  # ocb
            r.skip(2)  # flags
            r.skip(2)  # y_rot
            r.skip(2)  # box_num
            items.append(
                TR4LevelItem(
                    object_id=object_id,
                    x=x,
                    y=y,
                    z=z,
                    room_num=room_num,
                )
            )
    return items


def read_tr4_level(level_path: Path) -> TR4LevelData:
    data = level_path.read_bytes()
    r = _ByteReader(data)
    version = r.read_u32()
    if version != TR4_VERSION:
        raise ValueError(f"not a TR4 level file (version=0x{version:08X})")

    r.skip(6)  # room/object/bump texture page counts
    _read_chunk(r, "images32")
    _read_chunk(r, "images16")
    _read_chunk(r, "sky/font")

    level_data = _read_chunk(r, "level data")
    lr = _ByteReader(level_data)
    lr.skip(4)  # level number
    _skip_tr4_rooms(lr)
    _skip_tr4_object_meshes(lr)
    _skip_tr4_anims(lr)
    _skip_tr4_anim_data(lr)
    objects = _read_tr4_objects(lr)
    _skip_tr4_static_objects(lr)
    _skip_tr4_sprite_textures(lr)
    _skip_tr4_sprite_sequences(lr)
    _skip_tr4_cameras_and_sinks(lr)
    _skip_tr4_flyby_cameras(lr)
    _skip_tr4_sound_sources(lr)
    _skip_tr4_pathing_data(lr)
    _skip_tr4_animated_texture_ranges(lr)
    _skip_tr4_object_textures(lr)
    items = _read_tr4_items(lr)
    items.extend(_read_tr4_ai_items(lr))
    return TR4LevelData(objects=objects, items=items)


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)}: "
            f"{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 _get_zone_num(level_idx: int) -> int:
    if level_idx <= 12:
        return 0
    if level_idx <= 20:
        return 1
    if level_idx <= 26:
        return 2
    if level_idx <= 28:
        return 3
    return 4


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

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

    levels: list[dict[str, object]] = []
    for idx, level in enumerate(gameflow["levels"]):
        path = level.get("path")
        if not path:
            continue

        level_path = REPO_ROOT / "test/trx/games/tr4/levels" / path
        logger.info("Reading %s", level_path)
        parsed = read_tr4_level(level_path)
        title = strings["levels"][idx]["title"]
        levels.append(
            {
                "title": title,
                "zone_num": _get_zone_num(idx),
                "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()
