#!/usr/bin/env python3
"""Playthrough emulator worker: one PyBoy instance per awake coin, JSON lines over stdin / stdout.

Started by keepers/play-svc.mjs. Every request is one JSON object per line: {"id": n, "cmd": ..., ...}; every answer
is one line {"id": n, "ok": true, ...} or {"id": n, "ok": false, "error": ...}. While an action runs, the worker also
writes {"ev": "frame", "frame": n, "png": base64} lines (one per STREAM_EVERY emulated frames) that the svc paces out
to the WebSocket stream at ~5 fps.

Modes:
  real ROM   --rom play-emu/roms/pokered.gb (sha1 checked = Pokemon Red UE, file mode 600; never served)
  test ROM   --test-rom: PyBoy's bundled default_rom.gb; `poke` (fake milestone bytes) is allowed ONLY here and is
             part of the input log, so a replay reproduces it.

Commands:
  load {state: path|null}            load a savestate (null = the start state: start/<rom>.state if present, else boot)
  press {buttons:[..<=8], hold, wait} press each button for `hold` frames (1-30), release, run `wait` frames (0-600)
  walk {dir, tiles}                   = press [dir] x tiles with hold 16 / wait 4 (one tile per step)
  wait {frames}                       run frames (1-600)
  poke {addr, value}                  test ROM only
  ram {}                              RAM summary (map, x/y, party, badges, money, battle, walkable grid) + watched bytes
  frame {scale}                       current screen PNG (base64), scale 1-3 (nearest)
  save {path}                         write the savestate, return its sha256 (savestates stay private, the hash is public)
  hash {}                             sha256 of WRAM + screen (determinism checks)
Every action answer carries frame_start / frame_end (PyBoy frame counter since this worker started + the loaded
state's base) and the exact input list it executed (the public input log row).
"""
import argparse, base64, hashlib, io, json, os, sys

import pyboy as _pyboy
from pyboy import PyBoy

HERE = os.path.dirname(os.path.abspath(__file__))
REAL_SHA1 = "ea9bcae617fdf159b045185467ae58b2e4a48b9a"
BUTTONS = ("a", "b", "start", "select", "up", "down", "left", "right")
STREAM_EVERY = 12  # emulated frames between streamed frames (60 fps / 12 = 5 fps of game time)

# --- Pokemon Red RAM (pret/pokered symbols, RESEARCH-batch9 P2)
W_BADGES = 0xD356
W_CUR_MAP = 0xD35E
W_Y, W_X = 0xD361, 0xD362
W_PARTY_COUNT = 0xD163
W_PARTY_SPECIES = 0xD164
W_PARTY_MON1 = 0xD16B  # struct of 44 bytes: species +0, HP +1 (2 B BE), level +0x21, max HP +0x22 (2 B BE)
PARTY_STRUCT = 44
W_DEX_OWNED = 0xD2F7  # 19 bytes
W_DEX_SEEN = 0xD30A
W_MONEY = 0xD347  # 3 B BCD
W_HOF = 0xD5A2
W_TOWNS = 0xD70B  # 2 B, bit per town map id 0-10
W_IN_BATTLE = 0xD057
W_LAST_BLACKOUT = 0xD719
W_PLAYTIME = 0xDA41  # hours, maxed, minutes, seconds, frames

# milestone table: id -> (address, bit) for the bit flags (SPEC play)
BIT_MS = {i: (W_BADGES, i) for i in range(8)}
BIT_MS.update({
    10: (0xD74B, 5),  # got Pokedex
    11: (0xD74E, 1),  # got Oak's parcel
    12: (0xD7F2, 4),  # got SS ticket
    13: (0xD803, 0),  # got HM01
    14: (0xD76C, 2),  # got Poke Flute
    15: (0xD81B, 7),  # beat Rocket Hideout Giovanni
    16: (0xD838, 7),  # beat Silph Co Giovanni
})
for t in range(1, 11):  # towns 1-10 (0 = Pallet town, where everyone starts)
    BIT_MS[20 + t] = (W_TOWNS + (t // 8), t % 8)
CHAMPION_ID, CHAMPION_FLAG = 60, (0xD867, 1)
DEX_BASE = 40  # ids 40..54 = owned >= 10, 20, ... 150

MAP_NAMES = {0: "PALLET TOWN", 1: "VIRIDIAN CITY", 2: "PEWTER CITY", 3: "CERULEAN CITY", 4: "LAVENDER TOWN", 5: "VERMILION CITY",
             6: "CELADON CITY", 7: "FUCHSIA CITY", 8: "CINNABAR ISLAND", 9: "INDIGO PLATEAU", 10: "SAFFRON CITY", 12: "ROUTE 1",
             13: "ROUTE 2", 37: "RED'S HOUSE 1F", 38: "RED'S HOUSE 2F", 40: "OAK'S LAB", 41: "VIRIDIAN POKECENTER", 51: "VIRIDIAN FOREST",
             54: "PEWTER GYM", 59: "MT MOON 1F", 65: "CERULEAN GYM", 92: "VERMILION GYM", 134: "CELADON GYM", 157: "FUCHSIA GYM",
             166: "CINNABAR GYM", 178: "SAFFRON GYM", 45: "VIRIDIAN GYM", 118: "SS ANNE", 199: "ROCKET HIDEOUT B1F", 181: "SILPH CO 1F",
             245: "INDIGO PLATEAU LOBBY", 120: "HALL OF FAME"}


def sha1(path):
    h = hashlib.sha1()
    with open(path, "rb") as f:
        h.update(f.read())
    return h.hexdigest()


def popcount(bs):
    return sum(bin(b).count("1") for b in bs)


class Emu:
    def __init__(self, rom, test):
        self.test = test
        self.rom = rom
        self.pb = PyBoy(rom, window="null", sound_emulated=False)
        self.pb.set_emulation_speed(0)
        self.stream = None  # callback(frame_no, png_b64)
        self.base = 0  # frame counter of the loaded state (kept in the savestate side file)
        self.n = 0  # frames run since the load

    # ---------------------------------------------------------------- state
    def start_state_path(self):
        name = "test" if self.test else "pokered"
        return os.path.join(HERE, "start", name + ".state")

    def load(self, path):
        if path is None:
            p = self.start_state_path()
            if os.path.exists(p):
                path = p
        if path is None:
            # boot: deterministic power-on (no state file yet); the frame counter starts at 0
            self.pb.stop(False)
            self.pb = PyBoy(self.rom, window="null", sound_emulated=False)
            self.pb.set_emulation_speed(0)
            self.base = 0
            self.n = 0
            self._tick(1)
            return {"loaded": "boot", "frame": self.frame()}
        with open(path, "rb") as f:
            data = f.read()
        self.pb.load_state(io.BytesIO(data))
        meta = path + ".json"
        self.base = json.load(open(meta))["frame"] if os.path.exists(meta) else 0
        self.n = 0
        return {"loaded": os.path.basename(path), "state_sha256": hashlib.sha256(data).hexdigest(), "frame": self.frame()}

    def frame(self):
        return self.base + self.n

    def save(self, path):
        b = io.BytesIO()
        self.pb.save_state(b)
        data = b.getvalue()
        os.makedirs(os.path.dirname(path), exist_ok=True)
        tmp = path + ".tmp"
        with open(tmp, "wb") as f:
            f.write(data)
        os.chmod(tmp, 0o600)
        os.replace(tmp, path)
        with open(path + ".json", "w") as f:
            json.dump({"frame": self.frame()}, f)
        return {"path": path, "state_sha256": hashlib.sha256(data).hexdigest(), "bytes": len(data), "frame": self.frame()}

    # ---------------------------------------------------------------- input
    def _tick(self, n):
        for _ in range(n):
            self.pb.tick(1, True)
            self.n += 1
            if self.stream and self.frame() % STREAM_EVERY == 0:
                self.stream(self.frame(), self.png(1))

    def press(self, buttons, hold, wait):
        if not isinstance(buttons, list) or not 1 <= len(buttons) <= 8 or any(b not in BUTTONS for b in buttons):
            raise ValueError("buttons: 1-8 of " + "|".join(BUTTONS))
        hold = int(hold)
        wait = int(wait)
        if not 1 <= hold <= 30 or not 0 <= wait <= 600:
            raise ValueError("hold 1-30, wait 0-600")
        f0 = self.frame()
        inputs = []
        for b in buttons:
            self.pb.button_press(b)
            self._tick(hold)
            self.pb.button_release(b)
            self._tick(1 + wait)
            inputs.append({"b": b, "hold": hold, "wait": wait})
        return {"frame_start": f0, "frame_end": self.frame(), "inputs": inputs}

    def wait(self, frames):
        frames = int(frames)
        if not 1 <= frames <= 600:
            raise ValueError("frames 1-600")
        f0 = self.frame()
        self._tick(frames)
        return {"frame_start": f0, "frame_end": self.frame(), "inputs": [{"wait": frames}]}

    def poke(self, addr, value):
        if not self.test:
            raise ValueError("poke is refused with the real ROM")
        addr = int(addr)
        value = int(value)
        if not 0xC000 <= addr <= 0xDFFF or not 0 <= value <= 255:
            raise ValueError("WRAM only")
        f = self.frame()
        self.pb.memory[addr] = value
        return {"frame_start": f, "frame_end": f, "inputs": [{"poke": addr, "value": value}]}

    def apply(self, inp):
        """one input-log entry (replay)"""
        if "poke" in inp:
            self.poke(inp["poke"], inp["value"])
        elif "b" in inp:
            self.press([inp["b"]], inp["hold"], inp["wait"])
        elif "wait" in inp:
            self.wait(inp["wait"])
        else:
            raise ValueError("bad input " + json.dumps(inp))

    # ---------------------------------------------------------------- reads
    def m(self, a):
        return self.pb.memory[a]

    def watched(self):
        """the bytes the milestone detector diffs, by milestone id: [addr, value, bit or threshold, set?]"""
        out = {}
        for i, (a, bit) in BIT_MS.items():
            v = self.m(a)
            out[i] = [a, v, bit, bool(v >> bit & 1)]
        owned = popcount([self.m(W_DEX_OWNED + i) for i in range(19)])
        for k in range(15):
            out[DEX_BASE + k] = [W_DEX_OWNED, owned, 10 * (k + 1), owned >= 10 * (k + 1)]
        a, bit = CHAMPION_FLAG
        v = self.m(a)
        hof = self.m(W_HOF)
        out[CHAMPION_ID] = [a, v, bit, bool(v >> bit & 1) or hof > 0, hof]
        return out

    def ram(self):
        m = self.m
        n = min(m(W_PARTY_COUNT), 6)
        party = []
        for i in range(n):
            base = W_PARTY_MON1 + i * PARTY_STRUCT
            party.append({"species": m(base), "hp": m(base + 1) << 8 | m(base + 2), "level": m(base + 0x21), "max_hp": m(base + 0x22) << 8 | m(base + 0x23)})
        money = m(W_MONEY), m(W_MONEY + 1), m(W_MONEY + 2)
        money = int("%02x%02x%02x" % money) if all((x >> 4) < 10 and (x & 15) < 10 for x in money) else None
        cur = m(W_CUR_MAP)
        out = {
            "map": cur, "map_name": MAP_NAMES.get(cur, "MAP %d" % cur), "x": m(W_X), "y": m(W_Y),
            "badges": popcount([m(W_BADGES)]), "badge_bits": m(W_BADGES), "party": party, "money": money,
            "in_battle": m(W_IN_BATTLE), "dex_owned": popcount([m(W_DEX_OWNED + i) for i in range(19)]),
            "dex_seen": popcount([m(W_DEX_SEEN + i) for i in range(19)]), "hof": m(W_HOF), "last_blackout": m(W_LAST_BLACKOUT),
            "towns": m(W_TOWNS) | m(W_TOWNS + 1) << 8,
            "play_time": [m(W_PLAYTIME), m(W_PLAYTIME + 2), m(W_PLAYTIME + 3)],
            "frame": self.frame(), "test_rom": self.test,
        }
        try:
            if not self.test:
                gw = self.pb.game_wrapper
                grid = gw._get_screen_walkable_matrix()
                out["walkable"] = ["".join("." if c else "#" for c in row) for row in grid.tolist()]
        except Exception as e:  # noqa: BLE001
            out["walkable_error"] = str(e)[:80]
        out["watched"] = {str(k): v for k, v in self.watched().items()}
        return out

    def png(self, scale=1):
        from PIL import Image
        img = self.pb.screen.image.convert("RGB")
        if scale > 1:
            img = img.resize((160 * scale, 144 * scale), Image.NEAREST)
        b = io.BytesIO()
        img.save(b, "PNG")
        return base64.b64encode(b.getvalue()).decode()

    def hash(self):
        return hashlib.sha256(bytes(self.pb.memory[0xC000:0xE000]) + self.pb.screen.ndarray.tobytes()).hexdigest()


def open_rom(args):
    if args.test_rom:
        return os.path.join(os.path.dirname(_pyboy.__file__), "default_rom.gb"), True
    rom = args.rom or os.path.join(HERE, "roms", "pokered.gb")
    if not os.path.exists(rom):
        raise SystemExit("no ROM at %s (use --test-rom for the pipeline test)" % rom)
    st = os.stat(rom)
    if st.st_mode & 0o077:
        raise SystemExit("ROM file must be mode 600 (chmod 600 %s)" % rom)
    if sha1(rom) != REAL_SHA1 and not args.any_rom:
        raise SystemExit("ROM sha1 is not Pokemon Red (UE) " + REAL_SHA1)
    return rom, False


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("--no-stream", action="store_true")
    args = ap.parse_args()
    rom, test = open_rom(args)
    emu = Emu(rom, test)
    out = sys.stdout

    def write(o):
        out.write(json.dumps(o, separators=(",", ":")) + "\n")
        out.flush()

    if not args.no_stream:
        emu.stream = lambda f, png: write({"ev": "frame", "frame": f, "png": png})
    write({"ev": "ready", "test_rom": test, "pyboy": getattr(_pyboy, "__version__", "?")})
    for line in sys.stdin:
        line = line.strip()
        if not line:
            continue
        rid = None
        try:
            req = json.loads(line)
            rid = req.get("id")
            c = req.get("cmd")
            if c == "load":
                r = emu.load(req.get("state"))
            elif c == "press":
                r = emu.press(req.get("buttons"), req.get("hold", 4), req.get("wait", 8))
            elif c == "walk":
                d = req.get("dir")
                tiles = int(req.get("tiles", 1))
                if d not in ("up", "down", "left", "right") or not 1 <= tiles <= 8:
                    raise ValueError("walk: dir up|down|left|right, tiles 1-8")
                r = emu.press([d] * tiles, 16, 4)
            elif c == "wait":
                r = emu.wait(req.get("frames", 30))
            elif c == "poke":
                r = emu.poke(req.get("addr"), req.get("value"))
            elif c == "apply":
                f0 = emu.frame()
                for inp in req.get("inputs", []):
                    emu.apply(inp)
                r = {"frame_start": f0, "frame_end": emu.frame()}
            elif c == "ram":
                r = emu.ram()
            elif c == "frame":
                r = {"frame": emu.frame(), "png": emu.png(max(1, min(3, int(req.get("scale", 1)))))}
            elif c == "save":
                r = emu.save(req["path"])
            elif c == "hash":
                r = {"hash": emu.hash(), "frame": emu.frame()}
            elif c == "quit":
                write({"id": rid, "ok": True})
                break
            else:
                raise ValueError("unknown cmd %r" % c)
            r["id"] = rid
            r["ok"] = True
            write(r)
        except Exception as e:  # noqa: BLE001
            write({"id": rid, "ok": False, "error": str(e)[:300]})
    emu.pb.stop(False)


if __name__ == "__main__":
    main()
