#!/usr/bin/env python3
"""Summarize unclassified layer0/layer1 tile pairs by tileset."""
from __future__ import annotations

import argparse
import json
from collections import Counter, defaultdict
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 fallback_kind(tile: int, layer1: int, info: dict) -> str:
    return "pass" if tile > 0 else "block"


def summarize(
    maps: dict[str, dict],
    classes: dict[str, dict[str, list[int]]],
    limit_per_tileset: int = 32,
    example_limit: int = 5,
) -> list[dict]:
    counters: dict[str, Counter[tuple[int, int]]] = defaultdict(Counter)
    examples: dict[str, dict[tuple[int, int], list[dict]]] = defaultdict(lambda: defaultdict(list))

    for map_name, info in sorted(maps.items()):
        tileset = (info.get("layerTilesets") or [info.get("tileset")])[0]
        entry = classes.get(tileset, {})
        known_tiles = set(entry.get("pass", [])) | set(entry.get("block", []))
        known_pairs = {tuple(pair) for pair in entry.get("passPairs", [])} | {
            tuple(pair) for pair in entry.get("blockPairs", [])
        }
        for index, (tile, layer1) in enumerate(zip(info["layers"][0], info["layers"][1])):
            pair = (tile, layer1)
            if tile in known_tiles or pair in known_pairs:
                continue
            counters[tileset][pair] += 1
            if len(examples[tileset][pair]) < example_limit:
                examples[tileset][pair].append({
                    "map": map_name,
                    "x": index % info["width"],
                    "y": index // info["width"],
                })

    rows = []
    for tileset in sorted(counters):
        for (tile, layer1), count in counters[tileset].most_common(limit_per_tileset):
            rows.append({
                "tileset": tileset,
                "tile": tile,
                "layer1": layer1,
                "count": count,
                "fallbackKind": fallback_kind(tile, layer1, maps[examples[tileset][(tile, layer1)][0]["map"]]),
                "examples": examples[tileset][(tile, layer1)],
            })
    return rows


def markdown(rows: list[dict], web_prefix: str = "../web") -> str:
    lines = [
        "# Tile Pair Class Gaps",
        "",
        "Generated from `out/maps.js` and `data/tile_classes.json`.",
        "",
        "These are the highest-use unclassified `layer0/layer1` pairs. "
        "Use these rows when a `layer0` tile is mixed and should be classified with `passPairs` or `blockPairs`.",
        "",
        "| tileset | tile | layer1 | count | fallback | examples |",
        "| --- | ---: | ---: | ---: | --- | --- |",
    ]
    for row in rows:
        examples = ", ".join(
            f"[{point['map']} {point['x']},{point['y']}]({web_prefix}/game.html?map={point['map']}"
            f"&startTile={point['x']},{point['y']}&focusTile={point['x']},{point['y']}&collision=1&overview=1)"
            for point in row["examples"]
        )
        lines.append(
            f"| {row['tileset']} | {row['tile']} | {row['layer1']} | {row['count']} | "
            f"{row['fallbackKind']} | {examples} |"
        )
    lines.append("")
    return "\n".join(lines)


def rows_by_tileset(rows: list[dict]) -> dict[str, list[dict]]:
    grouped: dict[str, list[dict]] = {}
    for row in rows:
        grouped.setdefault(row["tileset"], []).append(row)
    return {tileset: grouped[tileset] for tileset in sorted(grouped)}


def write_tileset_pages(rows: list[dict], out_dir: Path) -> list[dict]:
    out_dir.mkdir(parents=True, exist_ok=True)
    index_rows = []
    for tileset, tileset_rows in rows_by_tileset(rows).items():
        path = out_dir / f"{tileset}.md"
        path.write_text(markdown(tileset_rows, "../../web"), encoding="utf-8")
        index_rows.append({
            "tileset": tileset,
            "rows": len(tileset_rows),
            "uses": sum(row["count"] for row in tileset_rows),
            "path": path,
        })
    return index_rows


def index_markdown(index_rows: list[dict], out_dir: Path) -> str:
    lines = [
        "# Tile Pair Class Gap Index",
        "",
        "Generated from `out/tile_pair_class_gaps.json`.",
        "",
        "| tileset | rows | uses | page |",
        "| --- | ---: | ---: | --- |",
    ]
    for row in index_rows:
        rel = row["path"].relative_to(out_dir)
        lines.append(f"| {row['tileset']} | {row['rows']} | {row['uses']} | [{rel.name}]({rel.as_posix()}) |")
    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_pair_class_gaps.json")
    parser.add_argument("--out", type=Path, default=OUT / "tile_pair_class_gaps.md")
    parser.add_argument("--split-dir", type=Path, default=OUT / "tile_pair_class_gaps")
    parser.add_argument("--limit-per-tileset", type=int, default=32)
    parser.add_argument("--example-limit", type=int, default=5)
    args = parser.parse_args()

    rows = summarize(load_maps(args.maps), load_tile_classes(args.classes), args.limit_per_tileset, args.example_limit)
    args.json_out.write_text(json.dumps(rows, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    args.out.write_text(markdown(rows), encoding="utf-8")
    index_rows = write_tileset_pages(rows, args.split_dir)
    (args.split_dir / "index.md").write_text(index_markdown(index_rows, args.split_dir), encoding="utf-8")
    print(f"wrote {len(rows)} tile pair class gap rows -> {args.out}")


if __name__ == "__main__":
    main()
