#!/usr/bin/env python3
"""Render CNS tile maps to PNG previews using the browser prototype metadata."""
from __future__ import annotations

import argparse
from pathlib import Path

from decode_cns import decompress_cns, indices_to_rgb, parse_image, write_png
from export_map_js import parse_map


ROOT = Path(__file__).resolve().parents[1]
EXTRACT_FLD = ROOT / "extract_fld"
TILE_INDEX_MODES = [
    "row40_zero",
    "page20x6quad_zero",
    "page20x6quad",
    "row40",
    "page20x12",
    "vflip40",
    "column40x12",
]

def load_tileset_rgb(name: str) -> tuple[int, int, bytes, bytes]:
    src = EXTRACT_FLD / f"{name}.cns"
    decoded = decompress_cns(src.read_bytes())
    width, height, palette, pixels, bpp = parse_image(decoded)
    rgb = indices_to_rgb(width, height, palette, pixels, bpp)
    return width, height, rgb, rgb[:3]


def draw_layer(
    canvas: bytearray,
    canvas_width: int,
    map_info: dict,
    layer: list[int],
    tileset: tuple[int, int, bytes, bytes],
    tile_index_mode: str,
    layer_index: int = 0,
) -> None:
    tile_size = map_info["tileSize"]
    map_width = map_info["width"]
    map_height = map_info["height"]
    tileset_width, tileset_height, tileset_rgb, transparent = tileset
    empty_layer1_tiles = set(map_info.get("emptyLayer1Tiles") or [])

    for map_y in range(map_height):
        for map_x in range(map_width):
            tile = layer[map_y * map_width + map_x]
            if tile == 0:
                continue
            if layer_index > 0 and tile in empty_layer1_tiles:
                continue

            source = source_xy(tile, tile_index_mode, tileset_width, tileset_height, tile_size, map_info)
            if source is None:
                continue
            src_x, src_y = source
            if src_y + tile_size > tileset_height:
                continue
            if src_x + tile_size > tileset_width:
                continue

            dst_x = map_x * tile_size
            dst_y = map_y * tile_size
            for y in range(tile_size):
                for x in range(tile_size):
                    src = ((src_y + y) * tileset_width + src_x + x) * 3
                    color = tileset_rgb[src : src + 3]
                    if color == transparent:
                        continue
                    dst = ((dst_y + y) * canvas_width + dst_x + x) * 3
                    canvas[dst : dst + 3] = color


def source_xy(
    tile: int,
    mode: str,
    tileset_width: int,
    tileset_height: int,
    tile_size: int,
    map_info: dict,
) -> tuple[int, int] | None:
    if mode.endswith("_zero"):
        mode = mode.removesuffix("_zero")
        index = tile
    else:
        index = tile - 1
    if index < 0:
        return None
    if mode == "row40":
        columns = map_info.get("tilesetColumns", tileset_width // tile_size)
        return (index % columns) * tile_size, (index // columns) * tile_size
    if mode == "page20x6quad":
        columns = 20
        rows = 6
        page_columns = 2
        page_tile_count = columns * rows
        page = index // page_tile_count
        local = index % page_tile_count
        return (
            ((page % page_columns) * columns + (local % columns)) * tile_size,
            ((page // page_columns) * rows + (local // columns)) * tile_size,
        )
    if mode == "page20x12":
        columns = 20
        rows = tileset_height // tile_size
        page_tile_count = columns * rows
        page = index // page_tile_count
        local = index % page_tile_count
        return (page * columns + (local % columns)) * tile_size, (local // columns) * tile_size
    if mode == "vflip40":
        columns = map_info.get("tilesetColumns", tileset_width // tile_size)
        rows = tileset_height // tile_size
        return (index % columns) * tile_size, (rows - 1 - (index // columns)) * tile_size
    if mode == "column40x12":
        rows = tileset_height // tile_size
        return (index // rows) * tile_size, (index % rows) * tile_size
    raise ValueError(f"unknown tile index mode {mode}")


def render_map(
    src: Path,
    dst: Path,
    visual_mode: str,
    tile_index_mode: str,
    tileset_names: list[str] | None = None,
) -> None:
    map_info = parse_map(src)
    if tileset_names:
        if len(tileset_names) == 1:
            tileset_names = tileset_names * len(map_info["layers"])
        map_info["layerTilesets"] = tileset_names[: len(map_info["layers"])]
    tile_size = map_info["tileSize"]
    width = map_info["width"] * tile_size
    height = map_info["height"] * tile_size
    canvas = bytearray((8, 8, 8)) * width * height
    tilesets = {
        name: load_tileset_rgb(name)
        for name in dict.fromkeys(map_info["layerTilesets"])
    }

    if visual_mode in {"layer0", "bothSceneTilesets", "bothFirstTileset", "layer1UnderSceneTilesets", "layer1UnderFirstTileset"}:
        draw_layer(
            canvas,
            width,
            map_info,
            map_info["layers"][0],
            tilesets[map_info["layerTilesets"][0]],
            tile_index_mode,
            0,
        )
    if visual_mode in {"layer1UnderSceneTilesets", "layer1UnderFirstTileset"}:
        draw_layer(
            canvas,
            width,
            map_info,
            map_info["layers"][1],
            tilesets[map_info["layerTilesets"][1 if visual_mode == "layer1UnderSceneTilesets" else 0]],
            tile_index_mode,
            1,
        )
    if visual_mode == "layer1":
        draw_layer(
            canvas,
            width,
            map_info,
            map_info["layers"][1],
            tilesets[map_info["layerTilesets"][0]],
            tile_index_mode,
            1,
        )
    elif visual_mode == "bothSceneTilesets":
        draw_layer(
            canvas,
            width,
            map_info,
            map_info["layers"][1],
            tilesets[map_info["layerTilesets"][1]],
            tile_index_mode,
            1,
        )
    elif visual_mode == "bothFirstTileset":
        draw_layer(
            canvas,
            width,
            map_info,
            map_info["layers"][1],
            tilesets[map_info["layerTilesets"][0]],
            tile_index_mode,
            1,
        )

    dst.parent.mkdir(parents=True, exist_ok=True)
    write_png(dst, width, height, bytes(canvas))
    print(
        f"{src} -> {dst} "
        f"({map_info['width']}x{map_info['height']} tiles, "
        f"tilesets={','.join(map_info['layerTilesets'])}, "
        f"visual={visual_mode}, tileIndex={tile_index_mode})"
    )


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("inputs", nargs="+", type=Path)
    parser.add_argument("-o", "--out-dir", type=Path, default=ROOT / "out" / "previews")
    parser.add_argument(
        "--visual-mode",
        default="layer0",
        choices=["layer0", "layer1", "bothSceneTilesets", "bothFirstTileset", "layer1UnderSceneTilesets", "layer1UnderFirstTileset"],
    )
    parser.add_argument("--tile-index-mode", default="row40_zero", choices=TILE_INDEX_MODES)
    parser.add_argument("--tilesets", nargs="+", help="override layer tilesets, without .cns")
    args = parser.parse_args()

    for src in args.inputs:
        suffix_parts = []
        if args.tilesets:
            suffix_parts.append("_".join(args.tilesets))
        if args.visual_mode != "layer0":
            suffix_parts.append(args.visual_mode)
        if args.tile_index_mode != "row40_zero":
            suffix_parts.append(args.tile_index_mode)
        suffix = "" if not suffix_parts else f"_{'_'.join(suffix_parts)}"
        render_map(
            src,
            args.out_dir / f"{src.stem}{suffix}.png",
            args.visual_mode,
            args.tile_index_mode,
            args.tilesets,
        )


if __name__ == "__main__":
    main()
