#!/usr/bin/env python3
"""Summarize script handler table entries used by save-selector leaf streams."""
from __future__ import annotations

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

from probe_exe_scene_tables import read_sections, va_to_offset


ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "out"
HANDLER_TABLE_VA = 0x00440720
DEFAULT_HANDLER_VA = 0x0040239F
HANDLER_TABLE_SCAN_COUNT = 96


def dword_at(exe: bytes, sections: list[dict], va: int) -> int | None:
    offset = va_to_offset(sections, va)
    if offset is None or offset + 4 > len(exe):
        return None
    return struct.unpack_from("<I", exe, offset)[0]


def section_name_for_va(sections: list[dict], va: int) -> str | None:
    for section in sections:
        start = section["va"]
        end = start + section["raw_size"]
        if start <= va < end:
            return section["name"]
    return None


def handler_for_opcode(exe: bytes, sections: list[dict], opcode: int) -> dict:
    entry_va = HANDLER_TABLE_VA + opcode * 4
    handler_va = dword_at(exe, sections, entry_va)
    handler_section = section_name_for_va(sections, handler_va) if handler_va is not None else None
    entry = {
        "opcode": opcode,
        "opcodeHex": f"0x{opcode:02x}",
        "entryVa": entry_va,
        "entryVaHex": f"0x{entry_va:08x}",
        "handlerVa": handler_va,
        "handlerVaHex": f"0x{handler_va:08x}" if handler_va is not None else None,
        "handlerSection": handler_section,
        "isDefaultHandler": handler_va == DEFAULT_HANDLER_VA,
        "isCodeHandler": handler_section == ".text",
    }
    if handler_va is not None and handler_section == ".text" and handler_va != DEFAULT_HANDLER_VA:
        entry["streamEffect"] = analyze_stream_effect(exe, sections, handler_va)
    return entry


def table_handler_vas(exe: bytes, sections: list[dict]) -> list[int]:
    handlers = []
    for opcode in range(HANDLER_TABLE_SCAN_COUNT):
        handler = dword_at(exe, sections, HANDLER_TABLE_VA + opcode * 4)
        if handler is not None and section_name_for_va(sections, handler) == ".text":
            handlers.append(handler)
    return sorted(set(handlers))


def handler_bytes(exe: bytes, sections: list[dict], handler_va: int, max_size: int = 0x300) -> bytes:
    offset = va_to_offset(sections, handler_va)
    if offset is None:
        return b""
    next_handlers = [va for va in table_handler_vas(exe, sections) if va > handler_va]
    stop_va = min(next_handlers) if next_handlers else handler_va + max_size
    size = max(0, min(max_size, stop_va - handler_va))
    code = exe[offset: offset + size]
    ret = code.find(b"\xc3")
    if ret >= 0:
        return code[: ret + 1]
    return code


def analyze_stream_effect(exe: bytes, sections: list[dict], handler_va: int) -> dict:
    code = handler_bytes(exe, sections, handler_va)
    advances = []
    reads = []
    stores = []
    can_jump_to_stream_plus4 = False
    for offset in range(len(code) - 3):
        absolute = handler_va + offset
        if code[offset: offset + 3] == b"\x83\x40\x40" and offset + 3 < len(code):
            advances.append({
                "handlerOffsetHex": f"0x{offset:02x}",
                "vaHex": f"0x{absolute:08x}",
                "bytes": code[offset + 3],
            })
        if code[offset: offset + 2] in {b"\x8a\x40", b"\x8a\x48", b"\x8b\x40"} and offset + 2 < len(code):
            stream_offset = code[offset + 2]
            if stream_offset <= 0x10:
                reads.append({
                    "handlerOffsetHex": f"0x{offset:02x}",
                    "vaHex": f"0x{absolute:08x}",
                    "streamOffset": stream_offset,
                    "streamOffsetHex": f"0x{stream_offset:02x}",
                })
        if code[offset: offset + 3] == b"\x66\x8b\x40" and offset + 3 < len(code):
            stream_offset = code[offset + 3]
            if stream_offset <= 0x10:
                reads.append({
                    "handlerOffsetHex": f"0x{offset:02x}",
                    "vaHex": f"0x{absolute:08x}",
                    "streamOffset": stream_offset,
                    "streamOffsetHex": f"0x{stream_offset:02x}",
                    "width": 2,
                })
        if code[offset: offset + 3] == b"\x89\x41\x40":
            stores.append({
                "handlerOffsetHex": f"0x{offset:02x}",
                "vaHex": f"0x{absolute:08x}",
                "kind": "storeEaxToContextStream",
            })
    for offset in range(len(code) - 8):
        if code[offset: offset + 3] == b"\x8b\x40\x04" and b"\x89\x41\x40" in code[offset + 3: offset + 18]:
            can_jump_to_stream_plus4 = True
    return {
        "scanBytes": len(code),
        "fixedAdvances": advances,
        "streamReads": reads,
        "streamStores": stores,
        "canJumpToDwordAtPlus4": can_jump_to_stream_plus4,
    }


def leaf_word_refs(leaf_rows: list[dict]) -> dict[int, list[dict]]:
    refs: dict[int, list[dict]] = {}
    for row in leaf_rows:
        streams = [
            ("leaf", row.get("leafPointerHex"), row.get("words") or []),
            ("nested", row.get("nestedPointerHex"), row.get("nestedWords") or []),
        ]
        for stream_kind, stream_va, words in streams:
            if not stream_va:
                continue
            for word in words:
                value_hex = word.get("valueHex")
                if not isinstance(value_hex, str):
                    continue
                value = int(value_hex, 16)
                opcode = value & 0xFF
                refs.setdefault(opcode, []).append(
                    {
                        "streamKind": stream_kind,
                        "streamVaHex": stream_va,
                        "wordIndex": word.get("index"),
                        "wordVaHex": word.get("vaHex"),
                        "valueHex": value_hex,
                        "leafPointerHex": row.get("leafPointerHex"),
                        "nestedPointerHex": row.get("nestedPointerHex"),
                        "source": row.get("source"),
                        "target": row.get("target"),
                        "cns": word.get("cns"),
                        "pointer": word.get("pointer") is True,
                    }
                )
    return refs


def build_summary(exe: bytes, leaf_rows: list[dict]) -> dict:
    sections = read_sections(exe)
    refs = leaf_word_refs(leaf_rows)
    opcodes = sorted(refs)
    entries = []
    for opcode in opcodes:
        entry = handler_for_opcode(exe, sections, opcode)
        entry["referenceCount"] = len(refs[opcode])
        entry["references"] = refs[opcode]
        entries.append(entry)
    return {
        "handlerTableVa": HANDLER_TABLE_VA,
        "handlerTableVaHex": f"0x{HANDLER_TABLE_VA:08x}",
        "defaultHandlerVa": DEFAULT_HANDLER_VA,
        "defaultHandlerVaHex": f"0x{DEFAULT_HANDLER_VA:08x}",
        "scope": "first-byte handler candidates from save-selector leaf stream dwords",
        "opcodeCount": len(entries),
        "codeHandlerCount": sum(1 for entry in entries if entry["isCodeHandler"]),
        "defaultHandlerCount": sum(1 for entry in entries if entry["isDefaultHandler"]),
        "entries": entries,
    }


def markdown(summary: dict) -> str:
    lines = [
        "# Script Handler Table Candidates",
        "",
        "These rows map the low byte of each save-selector leaf-stream dword to the script handler table.",
        "They are first-byte handler candidates only; they do not prove instruction length or operand layout.",
        "Rows whose resolved handler is not in `.text` are likely operand/pointer low bytes, not real opcodes.",
        "",
        f"Handler table: `{summary['handlerTableVaHex']}`.",
        f"Distinct first-byte candidates: {summary['opcodeCount']} "
        f"({summary['codeHandlerCount']} resolve into `.text`, {summary['defaultHandlerCount']} use the default handler).",
        "",
        "| opcode | handler entry | handler | section | stream effect | default | refs | example leaf words |",
        "| --- | --- | --- | --- | --- | --- | ---: | --- |",
    ]
    for entry in summary["entries"]:
        examples = ", ".join(
            f"{ref['streamVaHex']}[{ref['wordIndex']}]={ref['valueHex']}"
            for ref in entry["references"][:6]
        )
        if entry["referenceCount"] > 6:
            examples += ", ..."
        effect = stream_effect_text(entry.get("streamEffect") or {})
        lines.append(
            f"| `{entry['opcodeHex']}` | `{entry['entryVaHex']}` | `{entry.get('handlerVaHex') or '-'}` | "
            f"{entry.get('handlerSection') or '-'} | {effect} | {'yes' if entry.get('isDefaultHandler') else 'no'} | "
            f"{entry['referenceCount']} | {examples} |"
        )
    lines.append("")
    return "\n".join(lines)


def stream_effect_text(effect: dict) -> str:
    if not effect:
        return "-"
    parts = []
    advances = sorted({item.get("bytes") for item in effect.get("fixedAdvances", []) if item.get("bytes")})
    if advances:
        parts.append("advance " + "/".join(f"+{value}" for value in advances))
    reads = sorted({item.get("streamOffsetHex") for item in effect.get("streamReads", []) if item.get("streamOffsetHex")})
    if reads:
        parts.append("read " + ",".join(reads[:8]))
    if effect.get("canJumpToDwordAtPlus4"):
        parts.append("jump [stream+4]")
    if effect.get("streamStores"):
        parts.append("writes stream")
    return "; ".join(parts) or "-"


def html_page(summary: dict) -> str:
    body_rows = []
    for entry in summary["entries"]:
        examples = "<br>".join(
            html.escape(f"{ref['streamVaHex']}[{ref['wordIndex']}]={ref['valueHex']}")
            for ref in entry["references"][:10]
        )
        if entry["referenceCount"] > 10:
            examples += "<br>..."
        effect = stream_effect_text(entry.get("streamEffect") or {})
        body_rows.append(
            "<tr>"
            f"<td><code>{html.escape(entry['opcodeHex'])}</code></td>"
            f"<td><code>{html.escape(entry['entryVaHex'])}</code></td>"
            f"<td><code>{html.escape(entry.get('handlerVaHex') or '-')}</code></td>"
            f"<td>{html.escape(entry.get('handlerSection') or '-')}</td>"
            f"<td>{html.escape(effect)}</td>"
            f"<td>{'yes' if entry.get('isDefaultHandler') else 'no'}</td>"
            f"<td>{entry['referenceCount']}</td>"
            f"<td>{examples}</td>"
            "</tr>"
        )
    return "\n".join([
        "<!doctype html>",
        '<html lang="en">',
        "<head>",
        '  <meta charset="utf-8">',
        '  <meta name="viewport" content="width=device-width, initial-scale=1">',
        "  <title>Script Handler Table Candidates</title>",
        "  <style>",
        "    body { margin: 24px; background: #101010; color: #eee; font: 14px system-ui, sans-serif; }",
        "    table { border-collapse: collapse; width: 100%; }",
        "    th, td { border: 1px solid #333; padding: 6px 8px; vertical-align: top; }",
        "    th { background: #1d1d1d; position: sticky; top: 0; }",
        "    code { color: #f5d76e; }",
        "  </style>",
        "</head>",
        "<body>",
        "  <h1>Script Handler Table Candidates</h1>",
        "  <p>These rows map the low byte of each save-selector leaf-stream dword to the script handler table. They are first-byte handler candidates only; rows whose resolved handler is not in <code>.text</code> are likely operand/pointer low bytes.</p>",
        f"  <p>Handler table: <code>{html.escape(summary['handlerTableVaHex'])}</code>. Distinct candidates: {summary['opcodeCount']}; code handlers: {summary['codeHandlerCount']}; default handlers: {summary['defaultHandlerCount']}.</p>",
        "  <table>",
        "    <thead><tr><th>opcode</th><th>entry</th><th>handler</th><th>section</th><th>stream effect</th><th>default</th><th>refs</th><th>examples</th></tr></thead>",
        "    <tbody>",
        "\n".join(body_rows) or '<tr><td colspan="8">No candidates.</td></tr>',
        "    </tbody>",
        "  </table>",
        "</body>",
        "</html>",
        "",
    ])


def write_outputs(summary: dict, out_dir: Path = OUT) -> None:
    out_dir.mkdir(parents=True, exist_ok=True)
    (out_dir / "script_handler_table.json").write_text(
        json.dumps(summary, ensure_ascii=False, separators=(",", ":")) + "\n",
        encoding="utf-8",
    )


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--exe", type=Path, default=ROOT / "Hwanse2.exe")
    parser.add_argument("--leaf-streams", type=Path, default=OUT / "save_selector_leaf_streams.json")
    parser.add_argument("--out-dir", type=Path, default=OUT)
    args = parser.parse_args()
    summary = build_summary(
        args.exe.read_bytes(),
        json.loads(args.leaf_streams.read_text(encoding="utf-8")),
    )
    write_outputs(summary, args.out_dir)
    print(f"wrote {summary['opcodeCount']} script handler candidates -> {args.out_dir / 'script_handler_table.json'}")


if __name__ == "__main__":
    main()
