from __future__ import annotations

import json
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any


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

SELECTION_JSON = OUT / "battle_monster_action_selection_review.json"
FRAME_PROBE_JSON = OUT / "battle_monster_action_frame_probe.json"
SHARED_ACTION_JSON = OUT / "battle_monster_shared_action_effect_review.json"
OUTPUT_JSON = OUT / "battle_monster_action_catalog.json"


def load_json(path: Path) -> dict[str, Any]:
    return json.loads(path.read_text(encoding="utf-8"))


def compact_display_entry(entry: dict[str, Any] | None) -> dict[str, Any] | None:
    if not entry:
        return None

    def compact_sound(sound: dict[str, Any]) -> dict[str, Any]:
        compact = {
            "vaHex": sound.get("vaHex"),
            "wlkNo": sound.get("wlkNo"),
            "wlkFileIndex": sound.get("wlkFileIndex"),
            "normalWlkNo": sound.get("normalWlkNo"),
            "altWlkNo": sound.get("altWlkNo"),
            "mode": sound.get("mode"),
            "summary": sound.get("summary"),
        }
        if sound.get("effectArgsHex"):
            compact["effectArgsHex"] = sound.get("effectArgsHex")
        return compact

    sounds = []
    for sound in entry.get("sounds") or []:
        sounds.append(compact_sound(sound))
    result_sounds = []
    for sound in entry.get("resultSounds") or []:
        result_sounds.append(compact_sound(sound))
    effect_sounds = []
    for sound in entry.get("effectSounds") or []:
        effect_sounds.append(compact_sound(sound))
    wait_barriers = []
    for wait in entry.get("waitBarriers") or []:
        wait_barriers.append(
            {
                "vaHex": wait.get("vaHex"),
                "opcode": wait.get("opcode"),
                "mode": wait.get("mode"),
                "maskHex": wait.get("maskHex"),
                "bytes": wait.get("bytes"),
                "summary": wait.get("summary"),
            }
        )
    helper_calls = []
    for helper in entry.get("helperCalls") or []:
        helper_calls.append(
            {
                "vaHex": helper.get("vaHex"),
                "helperId": helper.get("helperId"),
                "helperIdHex": helper.get("helperIdHex"),
                "bytes": helper.get("bytes"),
                "summary": helper.get("summary"),
            }
        )
    actor_flags = []
    for flag in entry.get("actorFlags") or []:
        actor_flags.append(
            {
                "vaHex": flag.get("vaHex"),
                "mode": flag.get("mode"),
                "maskHex": flag.get("maskHex"),
                "bytes": flag.get("bytes"),
                "summary": flag.get("summary"),
            }
        )
    repeat_loops = []
    for loop in entry.get("repeatLoops") or []:
        repeat_loops.append(
            {
                "vaHex": loop.get("vaHex"),
                "targetVaHex": loop.get("targetVaHex"),
                "repeatCount": loop.get("repeatCount"),
                "summary": loop.get("summary"),
            }
        )
    movements = []
    for movement in entry.get("movements") or []:
        movements.append(
            {
                "vaHex": movement.get("vaHex"),
                "movementMode": movement.get("movementMode"),
                "motionMode": movement.get("motionMode"),
                "motionKind": movement.get("motionKind"),
                "selector": movement.get("selector"),
                "selectorHex": movement.get("selectorHex"),
                "divisor": movement.get("divisor"),
                "stepDivisor": movement.get("stepDivisor"),
                "targetRangePolicy": movement.get("targetRangePolicy"),
                "selectorMeaning": movement.get("selectorMeaning"),
                "kind1Formula": movement.get("kind1Formula"),
                "kind2Formula": movement.get("kind2Formula"),
                "handlerEvidence": movement.get("handlerEvidence"),
                "summary": movement.get("summary"),
            }
        )
    return {
        "tableIndex": entry.get("index"),
        "localSlot": entry.get("monsterActionSlotIfPhaseMinus0x0a"),
        "pointerVaHex": entry.get("pointerVaHex"),
        "stopReason": entry.get("stopReason"),
        "categories": entry.get("categories") or [],
        "frameSelectorSequence": entry.get("frameSelectorSequence") or [],
        "frames": entry.get("frames") or [],
        "sounds": sounds,
        "resultSounds": result_sounds,
        "effectSounds": effect_sounds,
        "waitBarriers": wait_barriers,
        "helperCalls": helper_calls,
        "repeatLoops": repeat_loops,
        "actorFlags": actor_flags,
        "movements": movements,
        "positionWrites": entry.get("positionWrites") or [],
    }


def action_visual_class(display: dict[str, Any] | None) -> str:
    if not display:
        return "unmatched-display-entry"
    cats = set(display.get("categories") or [])
    has_movement = bool(display.get("movements"))
    has_effect_sound = "effect-sound" in cats
    has_repeat = "repeat-loop" in cats
    frame_count = len(display.get("frameSelectorSequence") or [])
    if has_movement and has_effect_sound:
        return "movement-plus-effect"
    if has_movement:
        return "movement"
    if has_effect_sound:
        return "effect-sound"
    if has_repeat:
        return "repeat-frame"
    if frame_count > 2:
        return "multi-frame"
    if frame_count:
        return "simple-frame"
    return "no-frame"


def main() -> None:
    selection = load_json(SELECTION_JSON)
    frame_probe = load_json(FRAME_PROBE_JSON)
    shared = load_json(SHARED_ACTION_JSON) if SHARED_ACTION_JSON.exists() else {}

    producer = selection.get("descriptorScriptProducer") or {}
    action_rows = producer.get("actionRows") or []
    descriptor_rows = producer.get("descriptorRows") or []

    frame_assets_by_cns: dict[str, dict[str, Any]] = {
        row.get("cns"): row for row in frame_probe.get("assetRows") or [] if row.get("cns")
    }
    shared_by_id: dict[int, dict[str, Any]] = {
        row.get("skillId"): row for row in shared.get("rows") or [] if isinstance(row.get("skillId"), int)
    }

    rows: list[dict[str, Any]] = []
    by_enemy: dict[str, dict[str, Any]] = {}
    match_counts = Counter()
    visual_class_counts = Counter()
    unmatched_by_cns = Counter()
    unmatched_by_phase = Counter()
    table_denominators = Counter()

    for row in action_rows:
        cns = row.get("cns") or ""
        display_phase = row.get("displayPhase")
        visible_slot = row.get("visibleSlot")
        frame_asset = frame_assets_by_cns.get(cns)
        display_entry = None
        match_status = "matched"
        if not frame_asset:
            match_status = "missing-frame-asset"
            unmatched_by_cns[cns] += 1
        else:
            for entry in frame_asset.get("entries") or []:
                if entry.get("index") == display_phase:
                    display_entry = entry
                    break
            if display_entry is None:
                match_status = "missing-display-phase"
                unmatched_by_phase[f"{cns}:{display_phase}"] += 1

        display = compact_display_entry(display_entry)
        visual_class = action_visual_class(display)
        match_counts[match_status] += 1
        visual_class_counts[visual_class] += 1
        if row.get("weightDenominator") is not None:
            table_denominators[row.get("weightDenominator")] += 1

        shared_id = row.get("sharedActionId")
        shared_row = shared_by_id.get(shared_id, {})
        action = {
            "enemyName": row.get("enemyName"),
            "cns": cns,
            "actorTableId": row.get("actorTableId"),
            "actorTableIdHex": row.get("actorTableIdHex"),
            "enemyStatIndex": row.get("enemyStatIndex"),
            "descriptorVaHex": row.get("descriptorVaHex"),
            "scriptVaHex": row.get("scriptVaHex"),
            "choiceTableIndex": row.get("choiceTableIndex"),
            "choiceTableMarkerVaHex": row.get("choiceTableMarkerVaHex"),
            "choiceTableRngRange": row.get("choiceTableRngRange"),
            "choiceTableJumpCount": row.get("choiceTableJumpCount"),
            "choiceTableSemantics": row.get("choiceTableSemantics"),
            "targetVaHex": row.get("targetVaHex"),
            "bytesHex": row.get("bytesHex"),
            "pairOrder": row.get("pairOrder"),
            "variantIndex": row.get("variantIndex"),
            "variantCount": row.get("variantCount"),
            "branchVariantStatus": row.get("branchVariantStatus"),
            "branchOpcode": row.get("branchOpcode"),
            "branchOpcodeBytesHex": row.get("branchOpcodeBytesHex"),
            "branchFieldOffsetHex": row.get("branchFieldOffsetHex"),
            "branchFieldMeaning": row.get("branchFieldMeaning"),
            "branchComparator": row.get("branchComparator"),
            "branchCompareValueHex": row.get("branchCompareValueHex"),
            "branchConditionKind": row.get("branchConditionKind"),
            "variantCondition": row.get("variantCondition"),
            "branchResolution": row.get("branchResolution"),
            "branchResolutionActorTableIdHex": row.get("branchResolutionActorTableIdHex"),
            "branchResolutionMeaning": row.get("branchResolutionMeaning"),
            "sharedActionId": shared_id,
            "sharedActionIdHex": row.get("sharedActionIdHex"),
            "sharedActionName": row.get("sharedActionName") or shared_row.get("name") or "",
            "visibleSlot": visible_slot,
            "visibleSlotHex": row.get("visibleSlotHex"),
            "displayPhase": display_phase,
            "displayPhaseHex": row.get("displayPhaseHex"),
            "weightCount": row.get("weightCount"),
            "weightDenominator": row.get("weightDenominator"),
            "weightPercent": row.get("weightPercent"),
            "choiceWeightKind": row.get("choiceWeightKind"),
            "weightNote": row.get("weightNote"),
            "targetScope": row.get("targetScope"),
            "resultFamily": row.get("resultFamily"),
            "sharedActionSummary": row.get("sharedActionSummary") or shared_row.get("summary") or {},
            "displayMatchStatus": match_status,
            "visualClass": visual_class,
            "displayEntry": display,
        }
        rows.append(action)

        key = f"{row.get('actorTableIdHex')}:{row.get('enemyName')}:{cns}"
        enemy = by_enemy.setdefault(
            key,
            {
                "actorTableId": row.get("actorTableId"),
                "actorTableIdHex": row.get("actorTableIdHex"),
                "enemyStatIndex": row.get("enemyStatIndex"),
                "enemyName": row.get("enemyName"),
                "cns": cns,
                "descriptorVaHex": row.get("descriptorVaHex"),
                "actions": [],
            },
        )
        enemy["actions"].append(action)

    descriptor_without_frame_asset = []
    for descriptor in descriptor_rows:
        cns = descriptor.get("cns") or ""
        if cns and cns not in frame_assets_by_cns:
            descriptor_without_frame_asset.append(
                {
                    "enemyName": descriptor.get("enemyName"),
                    "cns": cns,
                    "actorTableIdHex": descriptor.get("actorTableIdHex"),
                    "descriptorVaHex": descriptor.get("descriptorVaHex"),
                    "reason": "not present in battle_monster_action_frame_probe assetRows; likely player/companion-style enemy asset or non-z* battle asset",
                }
            )

    enemies = list(by_enemy.values())

    def enemy_sort_key(item: dict[str, Any]) -> tuple[bool, int, str]:
        index = item.get("enemyStatIndex")
        return (index is None, index if index is not None else 9999, item.get("enemyName") or "")

    enemies.sort(key=enemy_sort_key)

    report = {
        "version": 1,
        "kind": "hwanse-battle-monster-action-catalog",
        "title": "Static monster action catalog",
        "status": "static-monster-action-selection-and-display-slot-joined",
        "runtimeUsed": False,
        "source": [
            str(SELECTION_JSON.relative_to(ROOT)),
            str(FRAME_PROBE_JSON.relative_to(ROOT)),
            str(SHARED_ACTION_JSON.relative_to(ROOT)),
        ],
        "summary": [
            "Descriptor script3 VM bytecode selects both shared action id (+0x59) and local visible slot (+0x5a).",
            "The visible slot is joined to the monster CNS display table through display phase = 0x0a + slot.",
            "Opcode 0x2b/0x0a confirms script3 choice tables as RNG jump tables; repeated pointers are exact table weights.",
            "Slot-first conditional blocks are resolved by active object +0x05 actor-table-id checks, so shared CNS descriptors can select actor-specific action ids.",
            "This report is static EXE analysis only; no Wine/runtime observation is used.",
        ],
        "metrics": {
            "actionRows": len(rows),
            "enemyRows": len(enemies),
            "frameAssetRows": len(frame_assets_by_cns),
            "matchedDisplayEntries": match_counts.get("matched", 0),
            "missingFrameAssetRows": match_counts.get("missing-frame-asset", 0),
            "missingDisplayPhaseRows": match_counts.get("missing-display-phase", 0),
            "descriptorWithoutFrameAssetCount": len(descriptor_without_frame_asset),
            "displayMatchStatusCounts": dict(sorted(match_counts.items())),
            "visualClassCounts": dict(sorted(visual_class_counts.items())),
            "weightDenominatorCounts": {str(k): v for k, v in sorted(table_denominators.items())},
        },
        "knownGaps": [
            "btl_atu/btl_rsu/btl_smu are player/companion-style enemy assets, so their display tables are recovered from descriptor-adjacent static tables rather than monster rect-table evidence.",
            "Display scripts provide frame/sound/motion evidence, but exact battle timing in milliseconds remains a separate frame-gate calibration problem.",
        ],
        "descriptorWithoutFrameAsset": descriptor_without_frame_asset,
        "unmatchedByCns": dict(sorted(unmatched_by_cns.items())),
        "unmatchedByPhase": dict(sorted(unmatched_by_phase.items())),
        "enemies": enemies,
        "actions": rows,
    }

    OUTPUT_JSON.write_text(
        json.dumps(report, ensure_ascii=False, separators=(",", ":")) + "\n",
        encoding="utf-8",
    )
    print(f"wrote {OUTPUT_JSON}")


if __name__ == "__main__":
    main()
