#!/usr/bin/env python3
"""Render frequently used map tiles with tile IDs and usage counts."""
from __future__ import annotations

import argparse
from collections import Counter
from pathlib import Path

from decode_cns import write_png
from export_map_js import parse_map
from render_map_preview import load_tileset_rgb


ROOT = Path(__file__).resolve().parents[1]
EXTRACT_FLD = ROOT / "extract_fld"
OUT = ROOT / "out"

DIGITS = {
    "0": ["111", "101", "101", "101", "111"],
    "1": ["010", "110", "010", "010", "111"],
    "2": ["111", "001", "111", "100", "111"],
    "3": ["111", "001", "111", "001", "111"],
    "4": ["101", "101", "111", "001", "001"],
    "5": ["111", "100", "111", "001", "111"],
    "6": ["111", "100", "111", "101", "111"],
    "7": ["111", "001", "010", "010", "010"],
    "8": ["111", "101", "111", "101", "111"],
    "9": ["111", "101", "111", "001", "111"],
}


def put_pixel(canvas: bytearray, width: int, x: int, y: int, color: tuple[int, int, int]) -> None:
    if x < 0 or y < 0:
        return
    offset = (y * width + x) * 3
    if offset < 0 or offset + 3 > len(canvas):
        return
    canvas[offset : offset + 3] = bytes(color)


def draw_text(canvas: bytearray, width: int, x: int, y: int, text: str, color: tuple[int, int, int]) -> None:
    cursor = x
    for char in text:
        if char == " ":
            cursor += 4
            continue
        glyph = DIGITS.get(char)
        if not glyph:
            cursor += 4
            continue
        for gy, row in enumerate(glyph):
            for gx, value in enumerate(row):
                if value == "1":
                    put_pixel(canvas, width, cursor + gx, y + gy, color)
        cursor += 4


def blit_scaled_tile(
    canvas: bytearray,
    canvas_width: int,
    dst_x: int,
    dst_y: int,
    tile_id: int,
    tileset: tuple[int, int, bytes, bytes],
    scale: int,
) -> None:
    tile_size = 16
    tileset_width, tileset_height, rgb, transparent = tileset
    columns = tileset_width // tile_size
    src_x = (tile_id % columns) * tile_size
    src_y = (tile_id // columns) * tile_size
    if src_y + tile_size > tileset_height:
        return
    for y in range(tile_size):
        for x in range(tile_size):
            src = ((src_y + y) * tileset_width + src_x + x) * 3
            color = rgb[src : src + 3]
            if color == transparent:
                continue
            for sy in range(scale):
                for sx in range(scale):
                    dst = ((dst_y + y * scale + sy) * canvas_width + dst_x + x * scale + sx) * 3
                    canvas[dst : dst + 3] = color


def usage_for_tileset(tileset_name: str, maps: list[Path]) -> Counter[int]:
    counter: Counter[int] = Counter()
    for src in maps:
        info = parse_map(src)
        for layer_index, layer in enumerate(info["layers"]):
            if info["layerTilesets"][layer_index] != tileset_name:
                continue
            counter.update(tile for tile in layer if tile)
    return counter


def usage_by_tileset_from_map_data(map_data: dict[str, dict]) -> dict[str, Counter[int]]:
    counters: dict[str, Counter[int]] = {}
    for info in map_data.values():
        tilesets = info.get("layerTilesets") or [info.get("tileset")]
        for layer_index, layer in enumerate(info["layers"]):
            tileset_name = tilesets[layer_index]
            counter = counters.setdefault(tileset_name, Counter())
            counter.update(tile for tile in layer if tile)
    return counters


def render_gallery_from_counter(tileset_name: str, counter: Counter[int], dst: Path, limit: int) -> None:
    items = counter.most_common(limit)
    if not items:
        raise ValueError(f"no map usage found for {tileset_name}")

    scale = 3
    cell_w = 72
    cell_h = 62
    columns = 8
    rows = (len(items) + columns - 1) // columns
    width = columns * cell_w
    height = rows * cell_h
    canvas = bytearray((18, 18, 18)) * width * height
    tileset = load_tileset_rgb(tileset_name)

    for index, (tile_id, count) in enumerate(items):
        cell_x = (index % columns) * cell_w
        cell_y = (index // columns) * cell_h
        blit_scaled_tile(canvas, width, cell_x + 12, cell_y + 2, tile_id, tileset, scale)
        draw_text(canvas, width, cell_x + 2, cell_y + 51, str(tile_id), (255, 255, 255))
        draw_text(canvas, width, cell_x + 34, cell_y + 51, str(count), (180, 220, 255))

    dst.parent.mkdir(parents=True, exist_ok=True)
    write_png(dst, width, height, bytes(canvas))
    print(f"wrote {len(items)} tile usage entries -> {dst}")


def render_gallery(tileset_name: str, dst: Path, limit: int) -> None:
    maps = sorted(EXTRACT_FLD.glob("map[0-9]*.cns"))
    counter = usage_for_tileset(tileset_name, maps)
    render_gallery_from_counter(tileset_name, counter, dst, limit)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("tileset", help="tileset name without .cns")
    parser.add_argument("--out", type=Path, default=OUT / "tile_usage")
    parser.add_argument("--limit", type=int, default=64)
    args = parser.parse_args()

    render_gallery(args.tileset, args.out / f"{args.tileset}_usage.png", args.limit)


if __name__ == "__main__":
    main()
