#!/usr/bin/env python3
"""Summarize EXE pointer references for save-selector frontier records."""
from __future__ import annotations

import argparse
import html
import json
import struct
from pathlib import Path

from probe_exe_scene_tables import read_sections, offset_to_va, va_to_offset


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


def parse_hex(value: str | None) -> int | None:
    return int(value, 16) if value else None


def section_for_offset(sections: list[dict], offset: int) -> dict | None:
    for section in sections:
        start = section["raw"]
        end = start + section["raw_size"]
        if start <= offset < end:
            return section
    return None


def find_refs(exe: bytes, sections: list[dict], value: int) -> list[dict]:
    refs = []
    needle = struct.pack("<I", value)
    search = 0
    while True:
        hit = exe.find(needle, search)
        if hit < 0:
            break
        search = hit + 1
        section = section_for_offset(sections, hit)
        ref_va = offset_to_va(sections, hit)
        if section and ref_va is not None:
            refs.append({
                "section": section["name"],
                "refVaHex": f"0x{ref_va:08x}",
                "fileOffsetHex": f"0x{hit:06x}",
                "targetVaHex": f"0x{value:08x}",
            })
    return refs


def selector_label(row: dict) -> str:
    return f"{row.get('group')}:{row.get('slot')}"


def build_selector_metadata_ref_index(selector_rows: list[dict] | None) -> dict[tuple[str, str], list[dict]]:
    index: dict[tuple[str, str], list[dict]] = {}
    for row in selector_rows or []:
        label = selector_label(row)
        metadata_refs = [
            (
                row.get("rowPointerVaHex"),
                row.get("rowPointerHex"),
                "selector-group-row-pointer-entry",
            ),
            (
                row.get("selectedPointerVaHex"),
                row.get("selectedPointerHex"),
                "selector-row-selected-root-entry",
            ),
        ]
        for dword in row.get("firstDwords") or []:
            if not dword.get("pointer"):
                continue
            metadata_refs.append((
                dword.get("vaHex"),
                dword.get("valueHex"),
                "selector-root-pointer-entry",
            ))
        for ref_va_hex, target_va_hex, role in metadata_refs:
            if not ref_va_hex or not target_va_hex:
                continue
            index.setdefault((ref_va_hex, target_va_hex), []).append({
                "selector": label,
                "role": role,
            })
    return index


def classify_ref(
    ref: dict,
    selector_metadata_index: dict[tuple[str, str], list[dict]],
    current_selectors: set[str],
) -> dict:
    metadata_matches = selector_metadata_index.get((ref.get("refVaHex"), ref.get("targetVaHex"))) or []
    if not metadata_matches:
        return {
            **ref,
            "selectorRefRole": "unclassified-pointer-ref",
            "selectorRefSelector": None,
            "selectorMetadataRef": False,
        }
    current_match = next(
        (match for match in metadata_matches if match.get("selector") in current_selectors),
        metadata_matches[0],
    )
    role = current_match["role"]
    if current_match.get("selector") in current_selectors:
        role = f"current-{role}"
    return {
        **ref,
        "selectorRefRole": role,
        "selectorRefSelector": current_match.get("selector"),
        "selectorMetadataRef": True,
    }


def dwords_at(exe: bytes, sections: list[dict], va: int, count: int) -> list[dict]:
    offset = va_to_offset(sections, va)
    if offset is None:
        return []
    rows = []
    for index in range(count):
        item_offset = offset + index * 4
        if item_offset + 4 > len(exe):
            break
        item_va = offset_to_va(sections, item_offset)
        if item_va is None:
            break
        value = struct.unpack_from("<I", exe, item_offset)[0]
        rows.append({
            "index": index,
            "vaHex": f"0x{item_va:08x}",
            "valueHex": f"0x{value:08x}",
            "u16Lo": value & 0xffff,
            "u16Hi": value >> 16,
        })
    return rows


def refs_for_path(
    exe: bytes,
    sections: list[dict],
    path_hex: list[str],
    selector_metadata_index: dict[tuple[str, str], list[dict]],
    current_selectors: set[str],
) -> list[dict]:
    result = []
    for value_hex in path_hex:
        value = parse_hex(value_hex)
        if value is None:
            continue
        refs = [
            classify_ref(ref, selector_metadata_index, current_selectors)
            for ref in find_refs(exe, sections, value)
        ]
        result.append({
            "valueHex": value_hex,
            "refs": refs,
            "selectorMetadataRefCount": sum(1 for ref in refs if ref.get("selectorMetadataRef")),
            "nonMetadataRefCount": sum(1 for ref in refs if not ref.get("selectorMetadataRef")),
        })
    return result


def matching_selector_paths(reference_rows: list[dict], label: str, source: str, target: str) -> list[dict]:
    matches = []
    for row in reference_rows:
        if row.get("label") != label or row.get("kind") != "fieldMap":
            continue
        if row.get("resource") not in {source, target}:
            continue
        matches.append({
            "resource": row.get("resource"),
            "pathHex": row.get("pathHex") or [],
            "refVaHex": row.get("refVaHex"),
            "sceneIdHex": row.get("sceneIdHex"),
            "tilesets": row.get("tilesets") or [],
        })
    return matches


def build_rows(
    exe: bytes,
    sections: list[dict],
    frontier_rows: list[dict],
    selector_references: list[dict],
    selector_rows: list[dict] | None = None,
) -> list[dict]:
    rows = []
    selector_metadata_index = build_selector_metadata_ref_index(selector_rows)
    for frontier in frontier_rows:
        source = frontier["source"]
        target = frontier["target"]
        selector_labels = frontier.get("selectors") or []
        current_selectors = set(selector_labels)
        selector_paths = [
            path
            for label in selector_labels
            for path in matching_selector_paths(selector_references, label, source, target)
        ]
        path_values = []
        for path in selector_paths:
            for value in path["pathHex"]:
                if value not in path_values:
                    path_values.append(value)
        scene_pair_rows = []
        for pair in frontier.get("scenePairs") or []:
            leaf = parse_hex(pair.get("leafPointerHex"))
            source_record = parse_hex(pair.get("sourceRecordVaHex"))
            target_record = parse_hex(pair.get("targetRecordVaHex"))
            scene_pair_rows.append({
                **pair,
                "leafRefs": find_refs(exe, sections, leaf) if leaf is not None else [],
                "sourceRecordRefs": find_refs(exe, sections, source_record) if source_record is not None else [],
                "targetRecordRefs": find_refs(exe, sections, target_record) if target_record is not None else [],
                "leafDwords": dwords_at(exe, sections, leaf, 8) if leaf is not None else [],
            })
        selector_path_refs = refs_for_path(
            exe,
            sections,
            path_values,
            selector_metadata_index,
            current_selectors,
        )
        current_root_ref = next(
            (
                item
                for item in selector_path_refs
                if item.get("valueHex")
                and any(
                    ref.get("selectorRefRole") == "current-selector-row-selected-root-entry"
                    for ref in item.get("refs") or []
                )
            ),
            {},
        )
        rows.append({
            "source": source,
            "target": target,
            "selectors": selector_labels,
            "selectorPaths": selector_paths,
            "selectorPathRefs": selector_path_refs,
            "currentRootRefIsSelectorRowEntry": bool(current_root_ref),
            "currentRootNonMetadataRefCount": int(current_root_ref.get("nonMetadataRefCount") or 0),
            "routePromotionStatus": "blocked",
            "scenePairs": scene_pair_rows,
        })
    return rows


def markdown(rows: list[dict]) -> str:
    lines = [
        "# Save Selector Xrefs",
        "",
        "Pointer-reference summary for save-selector frontier candidates. These are data-flow breadcrumbs, not confirmed gameplay transitions.",
        "",
        f"Frontier edges: {len(rows)}.",
        "",
    ]
    for row in rows:
        lines.extend([
            f"## {row['source']} -> {row['target']}",
            "",
            f"Selectors: {', '.join(row['selectors']) or '-'}.",
            f"Current root row-entry ref: {row.get('currentRootRefIsSelectorRowEntry')} "
            f"(non-metadata refs: {row.get('currentRootNonMetadataRefCount')}, "
            f"promotion: {row.get('routePromotionStatus')}).",
            "",
            "| resource | selector path | scene | tilesets |",
            "| --- | --- | --- | --- |",
        ])
        for path in row["selectorPaths"]:
            lines.append(
                f"| {path['resource']} | {' -> '.join(path['pathHex'])} | "
                f"{path.get('sceneIdHex') or '-'} | {', '.join(path.get('tilesets') or []) or '-'} |"
            )
        lines.extend([
            "",
            "| value | refs |",
            "| --- | --- |",
        ])
        for item in row["selectorPathRefs"]:
            refs = ", ".join(
                f"{ref['section']} {ref['refVaHex']} "
                f"({ref.get('selectorRefRole') or 'unclassified-pointer-ref'})"
                for ref in item["refs"]
            ) or "-"
            lines.append(f"| {item['valueHex']} | {refs} |")
        lines.extend([
            "",
            "| leaf | source record | target record | leaf refs | source record refs | target record refs |",
            "| --- | --- | --- | --- | --- | --- |",
        ])
        for pair in row["scenePairs"]:
            leaf_refs = ", ".join(f"{ref['section']} {ref['refVaHex']}" for ref in pair["leafRefs"]) or "-"
            source_refs = ", ".join(f"{ref['section']} {ref['refVaHex']}" for ref in pair["sourceRecordRefs"]) or "-"
            target_refs = ", ".join(f"{ref['section']} {ref['refVaHex']}" for ref in pair["targetRecordRefs"]) or "-"
            lines.append(
                f"| {pair.get('leafPointerHex')} | {pair.get('sourceRecordVaHex')} | {pair.get('targetRecordVaHex')} | "
                f"{leaf_refs} | {source_refs} | {target_refs} |"
            )
        lines.append("")
    return "\n".join(lines)


def html_page(rows: list[dict]) -> str:
    sections = []
    for row in rows:
        path_rows = []
        for path in row["selectorPaths"]:
            path_rows.append(
                "<tr>"
                f"<td>{html.escape(path['resource'])}</td>"
                f"<td>{html.escape(' -> '.join(path['pathHex']))}</td>"
                f"<td>{html.escape(path.get('sceneIdHex') or '-')}</td>"
                f"<td>{html.escape(', '.join(path.get('tilesets') or []) or '-')}</td>"
                "</tr>"
            )
        ref_rows = []
        for item in row["selectorPathRefs"]:
            refs = "<br>".join(
                html.escape(
                    f"{ref['section']} {ref['refVaHex']} "
                    f"({ref.get('selectorRefRole') or 'unclassified-pointer-ref'})"
                )
                for ref in item["refs"]
            ) or "-"
            ref_rows.append(f"<tr><td>{html.escape(item['valueHex'])}</td><td>{refs}</td></tr>")
        pair_rows = []
        for pair in row["scenePairs"]:
            leaf_refs = "<br>".join(html.escape(f"{ref['section']} {ref['refVaHex']}") for ref in pair["leafRefs"]) or "-"
            source_refs = "<br>".join(html.escape(f"{ref['section']} {ref['refVaHex']}") for ref in pair["sourceRecordRefs"]) or "-"
            target_refs = "<br>".join(html.escape(f"{ref['section']} {ref['refVaHex']}") for ref in pair["targetRecordRefs"]) or "-"
            pair_rows.append(
                "<tr>"
                f"<td>{html.escape(pair.get('leafPointerHex') or '-')}</td>"
                f"<td>{html.escape(pair.get('sourceRecordVaHex') or '-')}</td>"
                f"<td>{html.escape(pair.get('targetRecordVaHex') or '-')}</td>"
                f"<td>{leaf_refs}</td><td>{source_refs}</td><td>{target_refs}</td>"
                "</tr>"
            )
        sections.append("\n".join([
            f"<h2>{html.escape(row['source'])} -> {html.escape(row['target'])}</h2>",
            f"<p>Selectors: {html.escape(', '.join(row['selectors']) or '-')}.</p>",
            "<p>"
            f"Current root row-entry ref: {html.escape(str(row.get('currentRootRefIsSelectorRowEntry')))}; "
            f"non-metadata refs: {html.escape(str(row.get('currentRootNonMetadataRefCount')))}; "
            f"promotion: {html.escape(str(row.get('routePromotionStatus')))}."
            "</p>",
            "<h3>Selector Paths</h3>",
            "<table><thead><tr><th>resource</th><th>selector path</th><th>scene</th><th>tilesets</th></tr></thead><tbody>",
            "\n".join(path_rows) or '<tr><td colspan="4">-</td></tr>',
            "</tbody></table>",
            "<h3>Path References</h3>",
            "<table><thead><tr><th>value</th><th>refs</th></tr></thead><tbody>",
            "\n".join(ref_rows) or '<tr><td colspan="2">-</td></tr>',
            "</tbody></table>",
            "<h3>Scene Pair References</h3>",
            "<table><thead><tr><th>leaf</th><th>source record</th><th>target record</th><th>leaf refs</th><th>source record refs</th><th>target record refs</th></tr></thead><tbody>",
            "\n".join(pair_rows) or '<tr><td colspan="6">-</td></tr>',
            "</tbody></table>",
        ]))
    return "\n".join([
        "<!doctype html>",
        '<html lang="en">',
        "<head>",
        '  <meta charset="utf-8">',
        '  <meta name="viewport" content="width=device-width, initial-scale=1">',
        "  <title>Save Selector Xrefs</title>",
        "  <style>",
        "    body { margin: 24px; background: #101010; color: #eee; font: 14px system-ui, sans-serif; }",
        "    table { border-collapse: collapse; width: 100%; margin: 12px 0 24px; }",
        "    th, td { border: 1px solid #333; padding: 6px 8px; vertical-align: top; }",
        "    th { background: #1d1d1d; position: sticky; top: 0; }",
        "  </style>",
        "</head>",
        "<body>",
        "  <h1>Save Selector Xrefs</h1>",
        "  <p>Pointer-reference summary for save-selector frontier candidates. These are data-flow breadcrumbs, not confirmed gameplay transitions.</p>",
        *sections,
        "</body>",
        "</html>",
        "",
    ])


def write_outputs(rows: list[dict], out_dir: Path = OUT) -> None:
    out_dir.mkdir(parents=True, exist_ok=True)
    (out_dir / "save_selector_xrefs.json").write_text(
        json.dumps(rows, ensure_ascii=False, indent=2) + "\n",
        encoding="utf-8",
    )
    (out_dir / "save_selector_xrefs.html").write_text(html_page(rows), encoding="utf-8")


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--exe", type=Path, default=ROOT / "Hwanse2.exe")
    parser.add_argument("--frontier", type=Path, default=OUT / "save_selector_frontier.json")
    parser.add_argument("--selector-references", type=Path, default=OUT / "save_scene_selector_references.json")
    parser.add_argument("--selectors", type=Path, default=OUT / "save_scene_selectors.json")
    parser.add_argument("--out-dir", type=Path, default=OUT)
    args = parser.parse_args()

    exe = args.exe.read_bytes()
    sections = read_sections(exe)
    rows = build_rows(
        exe,
        sections,
        json.loads(args.frontier.read_text(encoding="utf-8")),
        json.loads(args.selector_references.read_text(encoding="utf-8")),
        json.loads(args.selectors.read_text(encoding="utf-8")),
    )
    write_outputs(rows, args.out_dir)
    print(f"wrote {len(rows)} save selector xref rows -> {args.out_dir / 'save_selector_xrefs.html'}")


if __name__ == "__main__":
    main()
