#!/usr/bin/env python3
"""Validate and export tile passability classes."""
from __future__ import annotations

import argparse
import json
from pathlib import Path

from decode_cns import decompress_cns, parse_image


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


def normalize(data: object) -> dict[str, dict[str, list[int]]]:
    if not isinstance(data, dict):
        raise ValueError("tile class data must be an object")
    normalized: dict[str, dict[str, list[int]]] = {}
    for tileset, entry in data.items():
        if not isinstance(tileset, str):
            raise ValueError("tileset keys must be strings")
        if not isinstance(entry, dict):
            raise ValueError(f"{tileset}: entry must be an object")
        result: dict[str, list] = {}
        for state in ["pass", "block"]:
            values = entry.get(state, [])
            if not isinstance(values, list):
                raise ValueError(f"{tileset}.{state} must be a list")
            tiles = sorted(set(values))
            if any(not isinstance(tile, int) or tile < 0 for tile in tiles):
                raise ValueError(f"{tileset}.{state} must contain non-negative integers")
            result[state] = tiles
        for state in ["passPairs", "blockPairs"]:
            values = entry.get(state, [])
            if not isinstance(values, list):
                raise ValueError(f"{tileset}.{state} must be a list")
            pairs = []
            for value in values:
                if (
                    not isinstance(value, list)
                    or len(value) != 2
                    or not all(isinstance(tile, int) and tile >= 0 for tile in value)
                ):
                    raise ValueError(f"{tileset}.{state} must contain [layer0, layer1] integer pairs")
                pairs.append((value[0], value[1]))
            result[state] = [list(pair) for pair in sorted(set(pairs))]
        overlap = set(result["pass"]) & set(result["block"])
        if overlap:
            sample = ", ".join(str(tile) for tile in sorted(overlap)[:8])
            raise ValueError(f"{tileset}: tiles classified as both pass and block: {sample}")
        pair_overlap = {tuple(pair) for pair in result["passPairs"]} & {tuple(pair) for pair in result["blockPairs"]}
        if pair_overlap:
            sample = ", ".join(f"{a}/{b}" for a, b in sorted(pair_overlap)[:8])
            raise ValueError(f"{tileset}: tile pairs classified as both pass and block: {sample}")
        normalized[tileset] = result
    return normalized


def load(path: Path) -> dict[str, dict[str, list[int]]]:
    return normalize(json.loads(path.read_text(encoding="utf-8")))


def tileset_capacity(src_dir: Path, tileset: str) -> int:
    src = src_dir / f"{tileset}.cns"
    if not src.exists():
        raise ValueError(f"{tileset}: missing tileset image {src}")
    decoded = decompress_cns(src.read_bytes())
    width, height, _palette, _pixels, _bpp = parse_image(decoded)
    if width % 16 or height % 16:
        raise ValueError(f"{tileset}: tileset dimensions {width}x{height} are not 16px aligned")
    return (width // 16) * (height // 16)


def validate_capacities(data: dict[str, dict[str, list[int]]], src_dir: Path = EXTRACT_FLD) -> None:
    capacities = {}
    for tileset, entry in data.items():
        capacity = capacities.setdefault(tileset, tileset_capacity(src_dir, tileset))
        for state in ["pass", "block"]:
            bad = [tile for tile in entry.get(state, []) if tile >= capacity]
            if bad:
                sample = ", ".join(str(tile) for tile in bad[:8])
                raise ValueError(f"{tileset}.{state}: tile IDs outside 0..{capacity - 1}: {sample}")
        for state in ["passPairs", "blockPairs"]:
            bad_pairs = [
                pair for pair in entry.get(state, [])
                if pair[0] >= capacity or pair[1] >= capacity
            ]
            if bad_pairs:
                sample = ", ".join(f"{pair[0]}/{pair[1]}" for pair in bad_pairs[:8])
                raise ValueError(f"{tileset}.{state}: tile pairs outside 0..{capacity - 1}: {sample}")


def write_outputs(data: dict[str, dict[str, list[int]]], out_dir: Path) -> None:
    out_dir.mkdir(parents=True, exist_ok=True)
    (out_dir / "tile_classes.json").write_text(
        json.dumps(data, ensure_ascii=False, indent=2) + "\n",
        encoding="utf-8",
    )
    (out_dir / "tile_classes.js").write_text(
        "window.HWANSE_TILE_CLASSES = "
        + json.dumps(data, ensure_ascii=False, separators=(",", ":"))
        + ";\n",
        encoding="utf-8",
    )


def merge(base: dict[str, dict[str, list[int]]], patch: dict[str, dict[str, list[int]]]) -> dict[str, dict[str, list[int]]]:
    merged = {
        tileset: {
            "pass": list(entry.get("pass", [])),
            "block": list(entry.get("block", [])),
            "passPairs": [list(pair) for pair in entry.get("passPairs", [])],
            "blockPairs": [list(pair) for pair in entry.get("blockPairs", [])],
        }
        for tileset, entry in base.items()
    }
    for tileset, entry in patch.items():
        target = merged.setdefault(tileset, {"pass": [], "block": [], "passPairs": [], "blockPairs": []})
        patch_pass = set(entry.get("pass", []))
        patch_block = set(entry.get("block", []))
        patch_pass_pairs = {tuple(pair) for pair in entry.get("passPairs", [])}
        patch_block_pairs = {tuple(pair) for pair in entry.get("blockPairs", [])}
        target["pass"] = sorted((set(target["pass"]) - patch_block) | patch_pass)
        target["block"] = sorted((set(target["block"]) - patch_pass) | patch_block)
        target["passPairs"] = [
            list(pair)
            for pair in sorted(({tuple(pair) for pair in target["passPairs"]} - patch_block_pairs) | patch_pass_pairs)
        ]
        target["blockPairs"] = [
            list(pair)
            for pair in sorted(({tuple(pair) for pair in target["blockPairs"]} - patch_pass_pairs) | patch_block_pairs)
        ]
    return normalize(merged)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--src", type=Path, default=DATA / "tile_classes.json")
    parser.add_argument("--out", type=Path, default=OUT)
    parser.add_argument("--merge", type=Path, help="merge a browser-exported tile class JSON patch into --src")
    parser.add_argument("--write-src", action="store_true", help="write merged data back to --src")
    parser.add_argument("--tileset-dir", type=Path, default=EXTRACT_FLD, help="directory containing map_* tileset CNS files")
    args = parser.parse_args()

    data = load(args.src)
    if args.merge:
        patch = normalize(json.loads(args.merge.read_text(encoding="utf-8")))
        data = merge(data, patch)
    validate_capacities(data, args.tileset_dir)
    if args.merge and args.write_src:
        args.src.write_text(json.dumps(data, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
        print(f"updated source tile classes -> {args.src}")
    write_outputs(data, args.out)
    print(f"wrote tile classes for {len(data)} tilesets -> {args.out}")


if __name__ == "__main__":
    main()
