#!/usr/bin/env python3
"""Summarize overwrite order for selectionBuffer[0x20] before the current frontier."""
from __future__ import annotations

import argparse
import html
import json
from pathlib import Path
from typing import Any


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


def load_json(path: Path, fallback: Any) -> Any:
    if not path.exists():
        return fallback
    return json.loads(path.read_text(encoding="utf-8"))


def va_int(value: str | None) -> int:
    return int(value or "0", 16)


def opcode(value_hex: str | None) -> int:
    return int(value_hex or "0", 16) & 0xFF


def stream_plus(value_hex: str | None, index: int) -> int:
    return (int(value_hex or "0", 16) >> (index * 8)) & 0xFF


def step_summary(step: dict) -> dict:
    return {
        "vaHex": step.get("vaHex"),
        "valueHex": step.get("valueHex"),
        "opcodeHex": step.get("opcodeHex"),
        "meaning": step.get("meaning") or step.get("stopReason") or "",
        "stopReason": step.get("stopReason"),
    }


def build_rows(selection_writers: dict, writer_paths: list[dict]) -> list[dict]:
    frontier_rows = selection_writers.get("currentFrontierRows") or []
    if not frontier_rows:
        return []
    first_reader = min(frontier_rows, key=lambda row: va_int(row.get("vaHex")))
    reader_va = va_int(first_reader.get("vaHex"))
    writers = selection_writers.get("currentFrontierRootWritersBeforeFirstReader") or []
    path_by_writer = {row.get("writerVaHex"): row for row in writer_paths}

    chain = []
    for writer in sorted(writers, key=lambda row: va_int(row.get("vaHex"))):
        writer_va = va_int(writer.get("vaHex"))
        path = path_by_writer.get(writer.get("vaHex")) or {}
        trace = path.get("trace") or []
        following = [
            step_summary(step)
            for step in trace
            if va_int(step.get("vaHex")) > writer_va and va_int(step.get("vaHex")) < reader_va
            and (step.get("meaning") or step.get("stopReason"))
        ]
        overwritten_by = [
            later.get("vaHex")
            for later in writers
            if va_int(later.get("vaHex")) > writer_va
        ]
        value_hex = writer.get("valueHex")
        operation = {
            0x12: "select active branch-state slot",
            0x13: "match party/runtime slot",
        }.get(opcode(value_hex), "unknown")
        stream1 = stream_plus(value_hex, 1)
        table = "primaryBranchState" if stream1 == 0 else "secondaryBranchState"
        chain.append({
            "writerVaHex": writer.get("vaHex"),
            "writerValueHex": value_hex,
            "streamPlus1Hex": f"0x{stream1:02x}",
            "tableChoice": table if opcode(value_hex) == 0x12 else "runtime/party matcher",
            "operation": operation,
            "distanceToReaderBytes": reader_va - writer_va,
            "overwrittenBy": overwritten_by,
            "isNearestLinearWriter": not overwritten_by,
            "followingGatesBeforeReader": [
                item for item in following
                if "gate" in item.get("meaning", "") or item.get("stopReason") in {"branch-or-fallthrough", "multiple-advances"}
            ],
        })

    nearest = next((row for row in reversed(chain) if row.get("isNearestLinearWriter")), None)
    return [{
        "source": "map1_01a",
        "target": "map2_02d",
        "selectionBufferOffsetHex": "0x20",
        "frontierReaderVaHex": first_reader.get("vaHex"),
        "frontierReaderValueHex": first_reader.get("valueHex"),
        "frontierCondition": (first_reader.get("frontierBranches") or [{}])[0].get("condition"),
        "writerChain": chain,
        "nearestLinearWriter": nearest,
        "conclusion": (
            "On the current selector root's linear address order, 0x005428bc is the nearest writer to "
            "selectionBuffer[0x20] before the 0x00542b0c frontier reader. Earlier 0x00542474, "
            "0x0054247c, and 0x00542484 writes are overwritten if execution reaches 0x005428bc. Promotion still needs "
            "control-flow proof that the path from 0x005428bc reaches 0x00542b0c and a strict tile hotspot."
        ),
    }]


def markdown(rows: list[dict]) -> str:
    lines = [
        "# Save Selector Writer Chain",
        "",
        "Overwrite order for `selectionBuffer[0x20]` before the current `map1_01a -> map2_02d` frontier reader.",
        "",
    ]
    for row in rows:
        lines.extend([
            f"## {row['source']} -> {row['target']}",
            "",
            f"- frontier reader: `{row.get('frontierReaderVaHex')}` `{row.get('frontierReaderValueHex')}`",
            f"- frontier condition: `{row.get('frontierCondition')}`",
            f"- conclusion: {row.get('conclusion')}",
            "",
            "| writer | value | operation | table/source | distance | overwritten by | gates before reader |",
            "| --- | --- | --- | --- | ---: | --- | --- |",
        ])
        for writer in row.get("writerChain") or []:
            gates = "<br>".join(
                f"{gate.get('vaHex')} {gate.get('opcodeHex')} {gate.get('meaning')}"
                for gate in writer.get("followingGatesBeforeReader") or []
            ) or "-"
            lines.append(
                f"| `{writer.get('writerVaHex')}` | `{writer.get('writerValueHex')}` | "
                f"{writer.get('operation')} | {writer.get('tableChoice')} | "
                f"{writer.get('distanceToReaderBytes')} | "
                f"{', '.join(writer.get('overwrittenBy') or []) or '-'} | {gates} |"
            )
        nearest = row.get("nearestLinearWriter") or {}
        lines.extend([
            "",
            f"Nearest linear writer: `{nearest.get('writerVaHex')}`.",
            "",
        ])
    return "\n".join(lines)


def html_page(rows: list[dict]) -> str:
    parts = [
        "<!doctype html><meta charset=\"utf-8\"><title>Save Selector Writer Chain</title>",
        "<style>body{font-family:system-ui,sans-serif;background:#111;color:#eee}table{border-collapse:collapse}td,th{border:1px solid #444;padding:6px 8px;vertical-align:top}code{color:#9bd4ff}</style>",
        "<h1>Save Selector Writer Chain</h1>",
        "<p>Overwrite order for <code>selectionBuffer[0x20]</code> before the current frontier reader.</p>",
    ]
    for row in rows:
        parts.extend([
            f"<h2>{html.escape(row['source'])} -&gt; {html.escape(row['target'])}</h2>",
            f"<p>frontier reader: <code>{html.escape(str(row.get('frontierReaderVaHex')))}</code> "
            f"<code>{html.escape(str(row.get('frontierReaderValueHex')))}</code></p>",
            f"<p>frontier condition: <code>{html.escape(str(row.get('frontierCondition')))}</code></p>",
            f"<p>{html.escape(str(row.get('conclusion')))}</p>",
            "<table><thead><tr><th>writer</th><th>value</th><th>operation</th><th>table/source</th><th>distance</th><th>overwritten by</th><th>gates before reader</th></tr></thead><tbody>",
        ])
        for writer in row.get("writerChain") or []:
            gates = "<br>".join(
                html.escape(f"{gate.get('vaHex')} {gate.get('opcodeHex')} {gate.get('meaning')}")
                for gate in writer.get("followingGatesBeforeReader") or []
            ) or "-"
            parts.append(
                "<tr>"
                f"<td><code>{html.escape(str(writer.get('writerVaHex')))}</code></td>"
                f"<td><code>{html.escape(str(writer.get('writerValueHex')))}</code></td>"
                f"<td>{html.escape(str(writer.get('operation')))}</td>"
                f"<td>{html.escape(str(writer.get('tableChoice')))}</td>"
                f"<td>{writer.get('distanceToReaderBytes')}</td>"
                f"<td>{html.escape(', '.join(writer.get('overwrittenBy') or []) or '-')}</td>"
                f"<td>{gates}</td>"
                "</tr>"
            )
        nearest = row.get("nearestLinearWriter") or {}
        parts.append("</tbody></table>")
        parts.append(f"<p>Nearest linear writer: <code>{html.escape(str(nearest.get('writerVaHex')))}</code>.</p>")
    return "\n".join(parts)


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


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--out-dir", type=Path, default=OUT)
    args = parser.parse_args()
    rows = build_rows(
        load_json(args.out_dir / "save_selector_selection_writers.json", {}),
        load_json(args.out_dir / "save_selector_current_writer_paths.json", []),
    )
    write_outputs(rows, args.out_dir)
    print(f"wrote {len(rows)} save selector writer chain rows -> {args.out_dir / 'save_selector_writer_chain.json'}")


if __name__ == "__main__":
    main()
