#!/usr/bin/env python3
"""Dump one extracted scene event record with nearby raw data."""
from __future__ import annotations

import argparse
import json
import struct
import sys
from pathlib import Path

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

from extract_scene_events import extract_events, read_tail_dwords
from probe_exe_scene_tables import c_string, find_cns_strings, read_sections, va_to_offset


def parse_int(text: str) -> int:
    return int(text, 0)


def format_dword(value: int, strings: dict[int, str]) -> str:
    lo = value & 0xFFFF
    hi = value >> 16
    if value in strings:
        return f"0x{value:08x} -> {strings[value]}"
    if 0x00400000 <= value <= 0x00600000:
        return f"0x{value:08x} -> ptr"
    if 0 <= lo < 512 and 0 <= hi < 512:
        return f"0x{value:08x} ({lo},{hi})"
    return f"0x{value:08x}"


def dump_raw_window(data: bytes, sections: list[dict], strings: dict[int, str], start_va: int, count: int) -> None:
    offset = va_to_offset(sections, start_va)
    if offset is None:
        print(f"raw @0x{start_va:08x}: not mapped")
        return
    print(f"raw dwords @0x{start_va:08x}:")
    for index in range(count):
        off = offset + index * 4
        if off + 4 > len(data):
            break
        va = start_va + index * 4
        value = struct.unpack_from("<I", data, off)[0]
        print(f"  0x{va:08x}: {format_dword(value, strings)}")


def choose_record(events: list[dict], record_va: int | None, map_name: str | None) -> dict:
    matches = events
    if record_va is not None:
        matches = [event for event in matches if event["recordVa"] == record_va]
    if map_name is not None:
        matches = [event for event in matches if event["map"] == map_name]
    if not matches:
        raise SystemExit("no matching scene event record")
    if len(matches) > 1:
        print("multiple records matched; pass --record-va to select one:")
        for event in matches:
            print(
                f"  0x{event['recordVa']:08x} {event['map']} "
                f"{event['sceneIdHex']} kind={event['eventKind']} "
                f"active={len(event.get('activePoints', event['points']))} "
                f"points={len(event['points'])}"
            )
        raise SystemExit(2)
    return matches[0]


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--exe", type=Path, default=Path("Hwanse2.exe"))
    parser.add_argument("--extract-dir", type=Path, default=Path("extract_fld"))
    parser.add_argument("--record-va", type=parse_int)
    parser.add_argument("--map")
    parser.add_argument("--raw-count", type=int, default=40)
    parser.add_argument("--json", action="store_true")
    args = parser.parse_args()

    data = args.exe.read_bytes()
    events = extract_events(data, args.extract_dir)
    record = choose_record(events, args.record_va, args.map)
    if args.json:
        print(json.dumps(record, ensure_ascii=False, indent=2))
        return

    sections = read_sections(data)
    strings = find_cns_strings(data, sections)
    print(
        f"{record['map']} {record['sceneIdHex']} kind={record['eventKind']} "
        f"record=0x{record['recordVa']:08x}"
    )
    print(
        f"active={len(record.get('activePoints', record['points']))} "
        f"inBounds={record.get('inBoundsPointCount', len(record['points']))} "
        f"points={len(record['points'])} pointTable=0x{record['pointTableVa']:08x} "
        f"raw={record.get('rawPointCount', len(record['points']))} "
        f"hint={record['pointCountHint']}"
    )
    print(
        f"dispatch=0x{record.get('eventDispatchVa', record['recordVa'] + 12):08x} "
        f"refs={len(record.get('eventDispatchRefs', []))}"
    )
    for ref in record.get("eventDispatchRefs", [])[:16]:
        print(
            f"  ref {ref['section']} va=0x{ref['refVa']:08x} "
            f"file=0x{ref['refFileOffset']:06x}"
        )
        if ref.get("conditionVa"):
            first = ref.get("conditionFirstDwords", [])
            first_text = ", ".join(item["hex"] for item in first[:4])
            print(f"    condition=0x{ref['conditionVa']:08x} first=[{first_text}]")
            if ref.get("conditionPayloadVa"):
                linked = ",".join(ref.get("conditionLinkedStrings", [])) or "-"
                print(f"    payload=0x{ref['conditionPayloadVa']:08x} links={linked}")
                payload = ref.get("conditionPayloadDwords", [])
                payload_text = ", ".join(item["hex"] for item in payload[:6])
                print(f"    payload first=[{payload_text}]")
    print("active points:")
    for point in record.get("activePoints", record["points"]):
        print(f"  ({point['x']},{point['y']})")
    print("in-bounds points:")
    for point in record.get("inBoundsPoints", record["points"]):
        print(f"  ({point['x']},{point['y']})")
    if record.get("rawPoints") and record.get("rawPoints") != record.get("inBoundsPoints", record["points"]):
        print("raw point table:")
        for point in record["rawPoints"]:
            mark = "*" if point in record.get("inBoundsPoints", record["points"]) else " "
            print(f" {mark}({point['x']},{point['y']})")

    print("tail dwords after point table:")
    for item in read_tail_dwords(data, sections, record["pointTableEndVa"], count=16):
        string = f" -> {item['string']}" if item.get("string") else ""
        print(
            f"  0x{item['va']:08x}: {item['hex']} "
            f"lo={item['lo']} hi={item['hi']}{string}"
        )

    print("pointer blocks:")
    for block in record["tailPointerBlocks"]:
        print(
            f"  {block['blockKind']} @0x{block['pointerVa']:08x} "
            f"first=0x{(block['firstValue'] or 0):x} ptrs={block['pointerCount']} "
            f"strings={','.join(block['strings']) or '-'}"
        )
        for item in block["firstDwords"][:10]:
            string = f" -> {item['string']}" if item.get("string") else ""
            print(f"    0x{item['va']:08x}: {item['hex']} lo={item['lo']} hi={item['hi']}{string}")

    dump_raw_window(data, sections, strings, record["recordVa"], args.raw_count)


if __name__ == "__main__":
    main()
