#!/usr/bin/env python3
"""Render tile-index mapping candidates for one CNS map."""
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"

MODES = [
    "row40_zero",
    "page20x6quad_zero",
    "page20x6quad",
    "row40",
    "page20x12",
    "vflip40",
    "column40x12",
]


def load_tileset(name: str) -> tuple[int, int, bytes, bytes]:
    decoded = decompress_cns((EXTRACT_FLD / f"{name}.cns").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 source_xy(tile: int, mode: str, tileset_width: int, tileset_height: int, tile_size: int) -> tuple[int, int]:
    if mode.endswith("_zero"):
        mode = mode.removesuffix("_zero")
        index = tile
    else:
        index = tile - 1
    if mode == "row40":
        columns = tileset_width // tile_size
        return (index % columns) * tile_size, (index // columns) * tile_size
    if mode == "page20x12":
        columns = 20
        page_tile_count = columns * (tileset_height // tile_size)
        page = index // page_tile_count
        local = index % page_tile_count
        return (page * columns + (local % columns)) * tile_size, (local // 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 == "vflip40":
        columns = 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 mode {mode}")


def draw_layer(
    canvas: bytearray,
    canvas_width: int,
    map_info: dict,
    layer: list[int],
    tileset: tuple[int, int, bytes, bytes],
    mode: str,
) -> None:
    tile_size = map_info["tileSize"]
    tileset_width, tileset_height, tileset_rgb, transparent = tileset
    for map_y in range(map_info["height"]):
        for map_x in range(map_info["width"]):
            tile = layer[map_y * map_info["width"] + map_x]
            if tile == 0:
                continue
            src_x, src_y = source_xy(tile, mode, tileset_width, tileset_height, tile_size)
            if src_x < 0 or src_y < 0 or src_x + tile_size > tileset_width or src_y + tile_size > tileset_height:
                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 render_candidate(map_info: dict, tilesets: list[tuple[int, int, bytes, bytes]], mode: str) -> tuple[int, int, bytes]:
    tile_size = map_info["tileSize"]
    width = map_info["width"] * tile_size
    height = map_info["height"] * tile_size
    canvas = bytearray((0, 0, 0)) * width * height
    for index, layer in enumerate(map_info["layers"]):
        draw_layer(canvas, width, map_info, layer, tilesets[min(index, len(tilesets) - 1)], mode)
    return width, height, bytes(canvas)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("map", type=Path)
    parser.add_argument("--tilesets", nargs="+", help="override layer tilesets, without .cns")
    parser.add_argument("--modes", nargs="+", default=MODES, choices=MODES)
    parser.add_argument("-o", "--out-dir", type=Path, default=ROOT / "out" / "mapping_probe")
    args = parser.parse_args()

    map_info = parse_map(args.map)
    names = args.tilesets or map_info["layerTilesets"]
    if len(names) == 1:
        names = names * len(map_info["layers"])
    tilesets = [load_tileset(name) for name in names]
    args.out_dir.mkdir(parents=True, exist_ok=True)

    for mode in args.modes:
        width, height, rgb = render_candidate(map_info, tilesets, mode)
        dst = args.out_dir / f"{args.map.stem}_{'_'.join(names)}_{mode}.png"
        write_png(dst, width, height, rgb)
        print(f"{mode}: {dst}")


if __name__ == "__main__":
    main()
