#!/usr/bin/env python3
"""Scan TR4 level files for item OCB values on the given object slots.

The walk through the level data mirrors the field order of
src/trx/game/level/format/format_tr4.c and the shared section readers.
"""

from __future__ import annotations

import argparse
import struct
import zlib
from pathlib import Path


class Reader:
    def __init__(self, data: bytes) -> None:
        self.data = data
        self.pos = 0

    def skip(self, count: int) -> None:
        self.pos += count

    def read(self, fmt: str) -> tuple:
        result = struct.unpack_from("<" + fmt, self.data, self.pos)
        self.pos += struct.calcsize("<" + fmt)
        return result

    def s8(self) -> int:
        return self.read("b")[0]

    def u8(self) -> int:
        return self.read("B")[0]

    def s16(self) -> int:
        return self.read("h")[0]

    def u16(self) -> int:
        return self.read("H")[0]

    def s32(self) -> int:
        return self.read("i")[0]

    def u32(self) -> int:
        return self.read("I")[0]


def read_chunk(reader: Reader) -> bytes:
    uncompressed_size = reader.u32()
    compressed_size = reader.u32()
    blob = reader.data[reader.pos : reader.pos + compressed_size]
    reader.skip(compressed_size)
    result = zlib.decompress(blob)
    assert len(result) == uncompressed_size
    return result


def skip_rooms(level: Reader) -> None:
    num_rooms = level.s16()
    for _ in range(num_rooms):
        level.skip(4 * 4)  # x, z, y_bottom, y_top
        num_data_words = level.u32()
        level.skip(2 * num_data_words)
        num_portals = level.s16()
        level.skip(32 * num_portals)
        num_z_sectors = level.s16()
        num_x_sectors = level.s16()
        level.skip(8 * num_z_sectors * num_x_sectors)
        level.skip(4)  # room colour
        num_lights = level.s16()
        level.skip(46 * num_lights)
        num_static_meshes = level.s16()
        level.skip(20 * num_static_meshes)
        level.skip(2 + 2)  # flipped room, flags
        level.skip(3)  # water scheme, reverb, alternate group

    floor_data_size = level.s32()
    level.skip(2 * floor_data_size)


def skip_to_items(level: Reader) -> None:
    level.skip(4)  # level number

    skip_rooms(level)

    # Mesh data + pointers.
    num_mesh_words = level.s32()
    level.skip(2 * num_mesh_words)
    num_mesh_pointers = level.s32()
    level.skip(4 * num_mesh_pointers)

    # Animation data.
    level.skip(40 * level.s32())  # animations
    level.skip(6 * level.s32())  # state changes
    level.skip(8 * level.s32())  # anim dispatches
    level.skip(2 * level.s32())  # anim commands
    level.skip(4 * level.s32())  # mesh trees (raw int32 count)
    level.skip(2 * level.s32())  # frames (raw word count)

    level.skip(18 * level.s32())  # moveables
    level.skip(32 * level.s32())  # static objects

    # Sprite textures + sequences.
    assert level.data[level.pos : level.pos + 3] == b"SPR"
    level.skip(3)
    level.skip(16 * level.s32())  # sprite textures
    level.skip(8 * level.s32())  # sprite sequences

    level.skip(16 * level.s32())  # cameras and sinks
    level.skip(40 * level.s32())  # flyby cameras
    level.skip(16 * level.s32())  # sound sources

    # Pathing data.
    num_boxes = level.s32()
    level.skip(8 * num_boxes)
    level.skip(2 * level.s32())  # overlaps
    level.skip(20 * num_boxes)  # zones

    # Animated textures.
    num_animated_words = level.s32()
    level.skip(2 * num_animated_words)
    level.skip(1)  # UV rotate range count

    # Object textures.
    assert level.data[level.pos : level.pos + 3] == b"TEX"
    level.skip(3)
    level.skip(38 * level.s32())


def scan_level(path: Path) -> list[tuple[int, int, int]]:
    reader = Reader(path.read_bytes())
    version = reader.u32()
    assert version == 0x00345254, f"not a TR4 level: {path}"

    reader.skip(3 * 2)  # texture page counts
    for _ in range(3):
        read_chunk(reader)  # images32, images16, sky/font
    level = Reader(read_chunk(reader))

    skip_to_items(level)

    results = []
    num_items = level.s32()
    for item_num in range(num_items):
        obj_id = level.s16()
        level.skip(2 + 12 + 2 + 2)  # room, position, y rotation, shade
        ocb = level.s16()
        level.skip(2)  # flags
        results.append((item_num, obj_id, ocb))
    return results


def parse_object_ids(values: list[str]) -> set[int]:
    result: set[int] = set()
    for value in values:
        for part in value.split(","):
            result.add(int(part))
    return result


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "paths", nargs="+", type=Path, help=".tr4 files or directories"
    )
    parser.add_argument(
        "-o",
        "--object-ids",
        action="append",
        default=[],
        metavar="ID[,ID...]",
        help="TOMB4 object slot number(s) to report; "
        "may be repeated or comma-separated. "
        "Without this option, every slot is reported.",
    )
    parser.add_argument(
        "--nonzero-ocb",
        action="store_true",
        help="only report items whose OCB is nonzero",
    )
    args = parser.parse_args()
    object_ids = parse_object_ids(args.object_ids)

    level_paths: list[Path] = []
    for path in args.paths:
        if path.is_dir():
            level_paths.extend(sorted(path.glob("*.tr4")))
        else:
            level_paths.append(path)

    for path in level_paths:
        try:
            items = scan_level(path)
        except Exception as exc:
            print(f"{path.name}: FAILED to parse ({exc})")
            continue

        for item_num, obj_id, ocb in items:
            if object_ids and obj_id not in object_ids:
                continue
            if args.nonzero_ocb and ocb == 0:
                continue
            print(f"{path.name}: item {item_num} slot={obj_id} ocb={ocb}")


if __name__ == "__main__":
    main()
