#!/usr/bin/env python3
"""Summarize how much of each map is covered by tile passability classes."""
from __future__ import annotations

import argparse
import json
from pathlib import Path

from summarize_map_tiles import load_maps
from tile_classes import load as load_tile_classes


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


def coverage_for_map(name: str, info: dict, classes: dict[str, dict[str, list[int]]]) -> dict:
    tileset = (info.get("layerTilesets") or [info.get("tileset")])[0]
    entry = classes.get(tileset, {})
    pass_tiles = set(entry.get("pass", []))
    block_tiles = set(entry.get("block", []))
    pass_pairs = {tuple(pair) for pair in entry.get("passPairs", [])}
    block_pairs = {tuple(pair) for pair in entry.get("blockPairs", [])}
    layer0 = info["layers"][0]
    layer1 = info["layers"][1]
    pass_count = 0
    block_count = 0
    fallback_pass = 0
    for tile, mask in zip(layer0, layer1):
        pair = (tile, mask)
        if pair in pass_pairs or tile in pass_tiles:
            pass_count += 1
        elif pair in block_pairs or tile in block_tiles:
            block_count += 1
        elif tile > 0:
            fallback_pass += 1
    fallback_block = len(layer0) - pass_count - block_count - fallback_pass
    classified = pass_count + block_count
    return {
        "map": name,
        "tileset": tileset,
        "tileCount": len(layer0),
        "classified": classified,
        "classifiedPct": round(classified * 100 / len(layer0), 2) if layer0 else 0,
        "pass": pass_count,
        "block": block_count,
        "fallbackPass": fallback_pass,
        "fallbackBlock": fallback_block,
    }


def summarize(maps: dict, classes: dict[str, dict[str, list[int]]]) -> list[dict]:
    return [
        coverage_for_map(name, info, classes)
        for name, info in sorted(maps.items())
    ]


def markdown(rows: list[dict]) -> str:
    lines = [
        "# Tile Class Coverage",
        "",
        "Generated from `out/maps.js` and `data/tile_classes.json`.",
        "",
        "`classified` is the number of cells whose `layer0` tile ID or `layer0/layer1` pair is explicitly in pass/block. "
        "`fallback` is still using the temporary runtime `layer0 > 0` heuristic.",
        "",
        "| map | tileset | classified | pass | block | fallback pass | fallback block |",
        "| --- | --- | ---: | ---: | ---: | ---: | ---: |",
    ]
    for row in rows:
        lines.append(
            f"| {row['map']} | {row['tileset']} | "
            f"{row['classified']}/{row['tileCount']} ({row['classifiedPct']}%) | "
            f"{row['pass']} | {row['block']} | "
            f"{row['fallbackPass']} | {row['fallbackBlock']} |"
        )
    lines.append("")
    return "\n".join(lines)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--maps", type=Path, default=OUT / "maps.js")
    parser.add_argument("--classes", type=Path, default=DATA / "tile_classes.json")
    parser.add_argument("--json-out", type=Path, default=OUT / "tile_class_coverage.json")
    parser.add_argument("--md-out", "--out", dest="md_out", type=Path, default=None)
    args = parser.parse_args()

    rows = summarize(load_maps(args.maps), load_tile_classes(args.classes))
    args.json_out.parent.mkdir(parents=True, exist_ok=True)
    args.json_out.write_text(json.dumps(rows, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    if args.md_out:
        args.md_out.write_text(markdown(rows), encoding="utf-8")
    print(f"wrote {len(rows)} tile class coverage rows -> {args.json_out}")


if __name__ == "__main__":
    main()
