#!/usr/bin/env python3
"""Replay a Playthrough coin from its public input log (anyone with their OWN ROM can run this).

PyBoy is deterministic (measured: same inputs -> same RAM, screen and savestate hash, also across save/load cycles),
so the start state + every input row of every run reproduce the coin's savestate hash and the frame where each
milestone's RAM bit flipped.

  play-emu/venv/bin/python play-emu/replay.py --rom pokered.gb  inputs/<mint>/run-*.jsonl  [--upto-run N] [--json]
  play-emu/venv/bin/python play-emu/replay.py --test-rom        inputs/<mint>/run-*.jsonl     (localnet pipeline test)

The svc serves the rows at /api/inputs/<mint>?run=N and the run records (start / end state hashes, input log hash) at
/api/runs/<mint>. Output: the state sha256 after each run, the frame of every milestone flip seen, and whether the row
frames were contiguous.
"""
import argparse, glob, hashlib, json, os, sys, tempfile

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from worker import Emu, open_rom  # noqa: E402


def canon(v):
    return json.dumps(v, sort_keys=True, separators=(",", ":"))


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--rom")
    ap.add_argument("--test-rom", action="store_true")
    ap.add_argument("--any-rom", action="store_true", help=argparse.SUPPRESS)
    ap.add_argument("--start", help="start savestate (default: the pad's start state for this ROM, else boot); 'boot' = power-on (replays start/pokered-intro.jsonl into the start state hash)")
    ap.add_argument("--upto-run", type=int)
    ap.add_argument("--json", action="store_true")
    ap.add_argument("--save-dir", help="keep the state after each run here (run-<n>.state + .json): a later replay starts from it with --start")
    ap.add_argument("files", nargs="+")
    a = ap.parse_args()
    rom, test = open_rom(a)
    emu = Emu(rom, test)
    if a.start == "boot":
        emu.start_state_path = lambda: os.devnull + ".none"
        a.start = None
    emu.load(a.start)
    rows = []
    for f in sorted(set(sum([glob.glob(x) for x in a.files], []))):
        with open(f) as fh:
            rows += [json.loads(l) for l in fh if l.strip()]
    rows.sort(key=lambda r: (r["run"], r["seq"]))
    out = {"runs": [], "flips": [], "contiguous": True}
    before = emu.watched()
    cur, run_rows = None, []
    tmp = a.save_dir or tempfile.mkdtemp()
    os.makedirs(tmp, mode=0o700, exist_ok=True)

    def close_run(run, rr):
        st = emu.save(os.path.join(tmp, "run-%d.state" % run))
        h = hashlib.sha256("\n".join(canon(r) for r in rr).encode()).hexdigest()
        out["runs"].append({"run": run, "state_sha256": st["state_sha256"], "frame": st["frame"], "input_log_hash": h, "rows": len(rr)})

    for r in rows:
        if a.upto_run is not None and r["run"] > a.upto_run:
            break
        if cur is not None and r["run"] != cur:
            close_run(cur, run_rows)
            run_rows = []
        cur = r["run"]
        if r["frame_start"] != emu.frame():
            out["contiguous"] = False
        for inp in r["inputs"]:
            emu.apply(inp)
        if r["frame_end"] != emu.frame():
            out["contiguous"] = False
        run_rows.append(r)
        now = emu.watched()
        for k, v in now.items():
            if v[3] and not before[k][3]:
                out["flips"].append({"id": k, "run": r["run"], "seq": r["seq"], "frame": emu.frame(), "evidence": v})
        before = now
    if cur is not None:
        close_run(cur, run_rows)
    if a.json:
        print(json.dumps(out))
    else:
        for x in out["runs"]:
            print("run %d: state %s frame %d inputs %s (%d rows)" % (x["run"], x["state_sha256"], x["frame"], x["input_log_hash"], x["rows"]))
        for f in out["flips"]:
            print("milestone %s flipped in run %d at frame %d" % (f["id"], f["run"], f["frame"]))
        print("frames contiguous:", out["contiguous"])


if __name__ == "__main__":
    main()
