#!/usr/bin/env python3
"""Summarize unclassified layer0 tile IDs by tileset for passability review."""
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 summarize(
    maps: dict[str, dict],
    classes: dict[str, dict[str, list[int]]],
    limit_per_tileset: int = 24,
    example_limit: int = 5,
) -> list[dict]:
    counters: dict[str, Counter[int]] = defaultdict(Counter)
    fallback_pass: dict[str, Counter[int]] = defaultdict(Counter)
    fallback_block: dict[str, Counter[int]] = defaultdict(Counter)
    examples: dict[str, dict[int, Counter[str]]] = defaultdict(lambda: defaultdict(Counter))

    for map_name, info in sorted(maps.items()):
        tileset = (info.get("layerTilesets") or [info.get("tileset")])[0]
        entry = classes.get(tileset, {})
        known = 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 tile, mask in zip(info["layers"][0], info["layers"][1]):
            if tile in known or (tile, mask) in known_pairs:
                continue
            counters[tileset][tile] += 1
            examples[tileset][tile][map_name] += 1
            if tile > 0:
                fallback_pass[tileset][tile] += 1
            else:
                fallback_block[tileset][tile] += 1

    rows = []
    for tileset in sorted(counters):
        for tile, count in counters[tileset].most_common(limit_per_tileset):
            map_examples = [
                {"map": name, "count": map_count}
                for name, map_count in examples[tileset][tile].most_common(example_limit)
            ]
            rows.append({
                "tileset": tileset,
                "tile": tile,
                "count": count,
                "fallbackPass": fallback_pass[tileset][tile],
                "fallbackBlock": fallback_block[tileset][tile],
                "fallbackKind": (
                    "mixed"
                    if fallback_pass[tileset][tile] and fallback_block[tileset][tile]
                    else "pass"
                    if fallback_pass[tileset][tile]
                    else "block"
                ),
                "examples": map_examples,
            })
    return rows


def fallback_kind(row: dict) -> str:
    if row.get("fallbackKind"):
        return row["fallbackKind"]
    if row["fallbackPass"] and row["fallbackBlock"]:
        return "mixed"
    if row["fallbackPass"]:
        return "pass"
    return "block"


def summary(rows: list[dict]) -> list[dict]:
    by_tileset: dict[str, dict] = {}
    for row in rows:
        item = by_tileset.setdefault(
            row["tileset"],
            {
                "tileset": row["tileset"],
                "rows": 0,
                "uses": 0,
                "passRows": 0,
                "blockRows": 0,
                "mixedRows": 0,
                "passUses": 0,
                "blockUses": 0,
                "mixedUses": 0,
            },
        )
        kind = fallback_kind(row)
        item["rows"] += 1
        item["uses"] += row["count"]
        item[f"{kind}Rows"] += 1
        item[f"{kind}Uses"] += row["count"]
    return [by_tileset[key] for key in sorted(by_tileset)]


def markdown(rows: list[dict]) -> str:
    summaries = summary(rows)
    lines = [
        "# Tile Class Gaps",
        "",
        "Generated from `out/maps.js` and `data/tile_classes.json`.",
        "",
        "These are the highest-use `layer0` tile IDs that are still not explicitly classified as pass/block. "
        "`fallback pass/block` shows how the temporary runtime `layer0 > 0` rule currently treats them. "
        "`mixed` rows need map-context review before promoting them to a layer0-only class.",
        "",
        "## Summary",
        "",
        "| tileset | rows | uses | pass rows/uses | block rows/uses | mixed rows/uses |",
        "| --- | ---: | ---: | ---: | ---: | ---: |",
    ]
    for item in summaries:
        lines.append(
            f"| {item['tileset']} | {item['rows']} | {item['uses']} | "
            f"{item['passRows']}/{item['passUses']} | "
            f"{item['blockRows']}/{item['blockUses']} | "
            f"{item['mixedRows']}/{item['mixedUses']} |"
        )
    lines.extend([
        "",
        "## Rows",
        "",
        "| tileset | tile | count | fallback | fallback pass | fallback block | review | example maps |",
        "| --- | ---: | ---: | --- | ---: | ---: | --- | --- |",
    ])
    for row in rows:
        examples = ", ".join(
            f"{item['map']}({item['count']})"
            for item in row["examples"]
        )
        review = f"../extract_fld/{row['tileset']}.cns"
        lines.append(
            f"| {row['tileset']} | {row['tile']} | {row['count']} | "
            f"{fallback_kind(row)} | "
            f"{row['fallbackPass']} | {row['fallbackBlock']} | "
            f"[CNS]({review}) | {examples} |"
        )
    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_gaps.json")
    parser.add_argument("--summary-json-out", type=Path, default=OUT / "tile_class_gap_summary.json")
    parser.add_argument("--out", type=Path, default=OUT / "tile_class_gaps.md")
    parser.add_argument("--limit-per-tileset", type=int, default=24)
    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.parent.mkdir(parents=True, exist_ok=True)
    args.json_out.write_text(json.dumps(rows, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    args.summary_json_out.write_text(json.dumps(summary(rows), ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    args.out.write_text(markdown(rows), encoding="utf-8")
    print(f"wrote {len(rows)} tile class gap rows -> {args.out}")


if __name__ == "__main__":
    main()
