#!/usr/bin/env python3
"""Probe Hwanse2.exe for tile-attribute table candidates.

Map tilesets are usually 40x12 16px tiles, so a per-tile attribute table could
be 480 bytes or 480 u16 values. This script scans PE data sections for compact
low-valued blocks and reports candidates that look like flags or small classes.
It is diagnostic only; a hit still needs code-reference or behavior analysis.
"""
from __future__ import annotations

import argparse
import json
import math
import struct
import sys
from collections import Counter
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))

from probe_exe_scene_tables import offset_to_va, read_sections


TILE_COUNT = 40 * 12


def entropy(counts: Counter[int], total: int) -> float:
    result = 0.0
    for count in counts.values():
        p = count / total
        result -= p * math.log2(p)
    return result


def text_ref_count(data: bytes, sections: list[dict], value: int) -> int:
    text = next(section for section in sections if section["name"] == ".text")
    raw = data[text["raw"] : text["raw"] + text["raw_size"]]
    needle = struct.pack("<I", value)
    count = 0
    search = 0
    while True:
        hit = raw.find(needle, search)
        if hit < 0:
            return count
        count += 1
        search = hit + 1


def nearby_ascii(data: bytes, offset: int, radius: int = 128) -> list[str]:
    start = max(0, offset - radius)
    end = min(len(data), offset + radius)
    items = []
    current = bytearray()
    for byte in data[start:end]:
        if 32 <= byte <= 126:
            current.append(byte)
            continue
        if len(current) >= 4:
            items.append(current.decode("ascii", "replace"))
        current.clear()
    if len(current) >= 4:
        items.append(current.decode("ascii", "replace"))
    return items[:6]


def candidate_score(counts: Counter[int], total: int, text_refs: int) -> float:
    largest = counts.most_common(1)[0][1] / total
    nonzero = 1 - (counts.get(0, 0) / total)
    balance = 1 - abs(nonzero - 0.35)
    return text_refs * 4 + balance * 2 + (1 - largest) + min(len(counts), 8) / 8


def scan_byte_tables(data: bytes, sections: list[dict], section_names: set[str], max_value: int) -> list[dict]:
    candidates = []
    for section in sections:
        if section["name"] not in section_names:
            continue
        raw_start = section["raw"]
        raw = data[raw_start : raw_start + section["raw_size"]]
        for local in range(0, len(raw) - TILE_COUNT + 1, 4):
            block = raw[local : local + TILE_COUNT]
            if not block:
                continue
            counts = Counter(block)
            if max(counts) > max_value or len(counts) < 2 or len(counts) > 12:
                continue
            zero_ratio = counts.get(0, 0) / TILE_COUNT
            if zero_ratio < 0.05 or zero_ratio > 0.95:
                continue
            file_offset = raw_start + local
            va = offset_to_va(sections, file_offset)
            if va is None:
                continue
            refs = text_ref_count(data, sections, va)
            candidates.append(
                {
                    "kind": "u8",
                    "section": section["name"],
                    "fileOffset": file_offset,
                    "va": va,
                    "textRefs": refs,
                    "unique": len(counts),
                    "zeroRatio": round(zero_ratio, 3),
                    "entropy": round(entropy(counts, TILE_COUNT), 3),
                    "counts": dict(sorted(counts.items())),
                    "nearbyAscii": nearby_ascii(data, file_offset),
                    "score": round(candidate_score(counts, TILE_COUNT, refs), 3),
                }
            )
    return candidates


def scan_u16_tables(data: bytes, sections: list[dict], section_names: set[str], max_value: int) -> list[dict]:
    candidates = []
    size = TILE_COUNT * 2
    for section in sections:
        if section["name"] not in section_names:
            continue
        raw_start = section["raw"]
        raw = data[raw_start : raw_start + section["raw_size"]]
        for local in range(0, len(raw) - size + 1, 4):
            values = [struct.unpack_from("<H", raw, local + index * 2)[0] for index in range(TILE_COUNT)]
            counts = Counter(values)
            if max(counts) > max_value or len(counts) < 2 or len(counts) > 12:
                continue
            zero_ratio = counts.get(0, 0) / TILE_COUNT
            if zero_ratio < 0.05 or zero_ratio > 0.95:
                continue
            file_offset = raw_start + local
            va = offset_to_va(sections, file_offset)
            if va is None:
                continue
            refs = text_ref_count(data, sections, va)
            candidates.append(
                {
                    "kind": "u16",
                    "section": section["name"],
                    "fileOffset": file_offset,
                    "va": va,
                    "textRefs": refs,
                    "unique": len(counts),
                    "zeroRatio": round(zero_ratio, 3),
                    "entropy": round(entropy(counts, TILE_COUNT), 3),
                    "counts": dict(sorted(counts.items())),
                    "nearbyAscii": nearby_ascii(data, file_offset),
                    "score": round(candidate_score(counts, TILE_COUNT, refs), 3),
                }
            )
    return candidates


def markdown(candidates: list[dict], limit: int) -> str:
    referenced = [item for item in candidates if item["textRefs"] > 0]
    lines = [
        "# Tile Attribute Table Candidates",
        "",
        "Generated by `tools/probe_tile_attribute_tables.py`.",
        "",
        "Candidates are 480-entry low-valued blocks, matching one 40x12 tileset page. "
        "This is a heuristic list, not proof of collision semantics.",
        "",
        f"Total candidates: {len(candidates)}. Candidates with direct `.text` refs: {len(referenced)}.",
        "",
        "| rank | kind | section | VA | text refs | unique | zero ratio | entropy | counts | nearby ascii |",
        "| ---: | --- | --- | --- | ---: | ---: | ---: | ---: | --- | --- |",
    ]
    for rank, item in enumerate(candidates[:limit], start=1):
        counts = ", ".join(f"{key}:{value}" for key, value in item["counts"].items())
        nearby = "<br>".join(item["nearbyAscii"]) or "-"
        lines.append(
            "| {rank} | {kind} | {section} | `{va}` | {textRefs} | {unique} | {zeroRatio} | {entropy} | {counts} | {nearby} |".format(
                rank=rank,
                kind=item["kind"],
                section=item["section"],
                va=f"0x{item['va']:08x}",
                textRefs=item["textRefs"],
                unique=item["unique"],
                zeroRatio=item["zeroRatio"],
                entropy=item["entropy"],
                counts=counts,
                nearby=nearby,
            )
        )
    lines.append("")
    return "\n".join(lines)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--exe", type=Path, default=Path("Hwanse2.exe"))
    parser.add_argument("--sections", default=".rdata,.data")
    parser.add_argument("--max-value", type=int, default=8)
    parser.add_argument("--limit", type=int, default=80)
    parser.add_argument("--json-out", type=Path, default=Path("out/tile_attribute_candidates.json"))
    parser.add_argument("--out", type=Path, default=Path("out/tile_attribute_candidates.md"))
    args = parser.parse_args()

    data = args.exe.read_bytes()
    sections = read_sections(data)
    section_names = set(args.sections.split(","))
    candidates = [
        *scan_byte_tables(data, sections, section_names, args.max_value),
        *scan_u16_tables(data, sections, section_names, args.max_value),
    ]
    candidates.sort(key=lambda item: item["score"], reverse=True)

    args.json_out.parent.mkdir(parents=True, exist_ok=True)
    args.json_out.write_text(json.dumps(candidates, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    args.out.parent.mkdir(parents=True, exist_ok=True)
    args.out.write_text(markdown(candidates, args.limit), encoding="utf-8")
    print(f"wrote {len(candidates)} tile attribute candidates -> {args.out}")
    print(f"wrote JSON -> {args.json_out}")


if __name__ == "__main__":
    main()
