#!/usr/bin/env python3
"""z80mini — a small Z80 interpreter for testing Spectrum machine-code
programs headlessly. Public domain, from thespeccy.com.

Covers the commonly used unprefixed opcodes plus the DD (IX) and ED
instructions our programs need. Flags: only Z and C are modelled, which
is enough for code that sticks to JR/RET/CALL cc with Z and C tests.

Hooks: on_out(port, value), on_in(port) -> value, on_halt() are
callables you may replace. HALT advances self.frame and continues.
"""

class Z80:
    def __init__(self):
        self.mem = bytearray(65536)
        self.pc = 0
        self.sp = 0
        self.ix = 0
        self.a = self.b = self.c = self.d = self.e = self.h = self.l = 0
        self.zf = self.cf = False
        self.frame = 0
        self.steps = 0
        self.on_out = lambda port, val: None
        self.on_in = lambda port: 0xFF
        self.on_halt = lambda: None
        self.call_stubs = {}          # addr -> callable, simulates CALL

    # -- register pair helpers --
    def hl(self): return (self.h << 8) | self.l
    def de(self): return (self.d << 8) | self.e
    def bc(self): return (self.b << 8) | self.c
    def set_hl(self, v): self.h, self.l = (v >> 8) & 0xFF, v & 0xFF
    def set_de(self, v): self.d, self.e = (v >> 8) & 0xFF, v & 0xFF
    def set_bc(self, v): self.b, self.c = (v >> 8) & 0xFF, v & 0xFF

    def f_byte(self):
        return (0x40 if self.zf else 0) | (0x01 if self.cf else 0)

    def set_f(self, v):
        self.zf, self.cf = bool(v & 0x40), bool(v & 0x01)

    def push(self, v):
        self.sp = (self.sp - 2) & 0xFFFF
        self.mem[self.sp] = v & 0xFF
        self.mem[(self.sp + 1) & 0xFFFF] = (v >> 8) & 0xFF

    def pop(self):
        v = self.mem[self.sp] | (self.mem[(self.sp + 1) & 0xFFFF] << 8)
        self.sp = (self.sp + 2) & 0xFFFF
        return v

    def imm8(self):
        v = self.mem[self.pc]; self.pc = (self.pc + 1) & 0xFFFF; return v

    def imm16(self):
        v = self.imm8(); return v | (self.imm8() << 8)

    def rel(self):
        d = self.imm8()
        return d - 256 if d > 127 else d

    REG = ['b', 'c', 'd', 'e', 'h', 'l', None, 'a']   # 6 = (HL)

    def get_r(self, i):
        if i == 6: return self.mem[self.hl()]
        return getattr(self, self.REG[i])

    def set_r(self, i, v):
        if i == 6: self.mem[self.hl()] = v
        else: setattr(self, self.REG[i], v)

    def get_rr(self, i):  # BC DE HL SP
        return [self.bc, self.de, self.hl, lambda: self.sp][i]()

    def set_rr(self, i, v):
        [self.set_bc, self.set_de, self.set_hl,
         lambda x: setattr(self, 'sp', x)][i](v)

    def alu(self, op, val):
        a = self.a
        if op == 0:                       # ADD
            r = a + val; self.cf = r > 255
        elif op == 1:                     # ADC
            r = a + val + self.cf; self.cf = r > 255
        elif op == 2:                     # SUB
            r = a - val; self.cf = r < 0
        elif op == 3:                     # SBC
            r = a - val - self.cf; self.cf = r < 0
        elif op == 4:                     # AND
            r = a & val; self.cf = False
        elif op == 5:                     # XOR
            r = a ^ val; self.cf = False
        elif op == 6:                     # OR
            r = a | val; self.cf = False
        else:                             # CP
            r = a - val; self.cf = r < 0
            self.zf = (r & 0xFF) == 0
            return
        self.a = r & 0xFF
        self.zf = self.a == 0

    def cond(self, i):  # NZ Z NC C
        return [not self.zf, self.zf, not self.cf, self.cf][i]

    def step(self):
        self.steps += 1
        op = self.imm8()

        if op == 0x00: return                                # NOP
        if op == 0x76: self.frame += 1; self.on_halt(); return
        if op in (0xF3, 0xFB): return                        # DI/EI

        # 8-bit loads 0x40-0x7F
        if 0x40 <= op <= 0x7F:
            self.set_r((op >> 3) & 7, self.get_r(op & 7)); return
        # ALU 0x80-0xBF
        if 0x80 <= op <= 0xBF:
            self.alu((op >> 3) & 7, self.get_r(op & 7)); return
        # ALU with immediate: C6 CE D6 DE E6 EE F6 FE
        if (op & 0xC7) == 0xC6:
            self.alu((op >> 3) & 7, self.imm8()); return

        top = op & 0xC7
        if top == 0x06: self.set_r((op >> 3) & 7, self.imm8()); return   # LD r,n
        if top == 0x04:                                                  # INC r
            i = (op >> 3) & 7
            v = (self.get_r(i) + 1) & 0xFF
            self.set_r(i, v); self.zf = v == 0; return
        if top == 0x05:                                                  # DEC r
            i = (op >> 3) & 7
            v = (self.get_r(i) - 1) & 0xFF
            self.set_r(i, v); self.zf = v == 0; return

        pair = op & 0xCF
        if pair == 0x01: self.set_rr((op >> 4) & 3, self.imm16()); return   # LD rr,nn
        if pair == 0x03: i = (op >> 4) & 3; self.set_rr(i, (self.get_rr(i) + 1) & 0xFFFF); return
        if pair == 0x0B: i = (op >> 4) & 3; self.set_rr(i, (self.get_rr(i) - 1) & 0xFFFF); return
        if pair == 0x09:                                                  # ADD HL,rr
            r = self.hl() + self.get_rr((op >> 4) & 3)
            self.cf = r > 0xFFFF; self.set_hl(r & 0xFFFF); return
        if pair == 0xC5:                                                  # PUSH
            i = (op >> 4) & 3
            v = [self.bc(), self.de(), self.hl(),
                 (self.a << 8) | self.f_byte()][i]
            self.push(v); return
        if pair == 0xC1:                                                  # POP
            i = (op >> 4) & 3
            v = self.pop()
            if i == 3: self.a = v >> 8; self.set_f(v & 0xFF)
            else: self.set_rr(i, v)
            return

        if op == 0x0A: self.a = self.mem[self.bc()]; return
        if op == 0x1A: self.a = self.mem[self.de()]; return
        if op == 0x02: self.mem[self.bc()] = self.a; return
        if op == 0x12: self.mem[self.de()] = self.a; return
        if op == 0x3A: self.a = self.mem[self.imm16()]; return
        if op == 0x32: self.mem[self.imm16()] = self.a; return
        if op == 0x2A:
            n = self.imm16()
            self.set_hl(self.mem[n] | (self.mem[(n + 1) & 0xFFFF] << 8)); return
        if op == 0x22:
            n = self.imm16()
            self.mem[n] = self.l; self.mem[(n + 1) & 0xFFFF] = self.h; return

        if op == 0x07:                                        # RLCA
            self.cf = bool(self.a & 0x80)
            self.a = ((self.a << 1) | (self.a >> 7)) & 0xFF; return
        if op == 0x0F:                                        # RRCA
            self.cf = bool(self.a & 1)
            self.a = ((self.a >> 1) | (self.a << 7)) & 0xFF; return
        if op == 0x17:                                        # RLA
            c = self.a >> 7
            self.a = ((self.a << 1) | (1 if self.cf else 0)) & 0xFF
            self.cf = bool(c); return
        if op == 0x1F:                                        # RRA
            c = self.a & 1
            self.a = (self.a >> 1) | (0x80 if self.cf else 0)
            self.cf = bool(c); return
        if op == 0x2F: self.a ^= 0xFF; return                 # CPL
        if op == 0x37: self.cf = True; return                 # SCF
        if op == 0x3F: self.cf = not self.cf; return          # CCF
        if op == 0xEB:                                        # EX DE,HL
            self.d, self.h = self.h, self.d
            self.e, self.l = self.l, self.e; return

        if op == 0x18: self._jr(True); return
        if op in (0x20, 0x28, 0x30, 0x38):
            self._jr(self.cond((op >> 3) & 3)); return
        if op == 0x10:                                        # DJNZ
            d = self.rel()
            self.b = (self.b - 1) & 0xFF
            if self.b: self.pc = (self.pc + d) & 0xFFFF
            return
        if op == 0xC3: self.pc = self.imm16(); return
        if (op & 0xC7) == 0xC2:                               # JP cc
            n = self.imm16()
            if self.cond((op >> 3) & 3): self.pc = n
            return
        if op == 0xE9: self.pc = self.hl(); return            # JP (HL)
        if op == 0xCD: self._call(self.imm16()); return
        if (op & 0xC7) == 0xC4:                               # CALL cc
            n = self.imm16()
            if self.cond((op >> 3) & 3): self._call(n)
            return
        if op == 0xC9: self.pc = self.pop(); return
        if (op & 0xC7) == 0xC0:                               # RET cc
            if self.cond((op >> 3) & 3): self.pc = self.pop()
            return

        if op == 0xD3: self.on_out(self.imm8(), self.a); return
        if op == 0xDB:
            port = (self.a << 8) | self.imm8()
            self.a = self.on_in(port) & 0xFF; return

        if op == 0xED:
            op2 = self.imm8()
            if op2 == 0xB0:                                   # LDIR
                while True:
                    self.mem[self.de()] = self.mem[self.hl()]
                    self.set_hl((self.hl() + 1) & 0xFFFF)
                    self.set_de((self.de() + 1) & 0xFFFF)
                    self.set_bc((self.bc() - 1) & 0xFFFF)
                    if self.bc() == 0: break
                return
            if op2 == 0x52:                                   # SBC HL,DE
                r = self.hl() - self.de() - (1 if self.cf else 0)
                self.cf = r < 0
                self.set_hl(r & 0xFFFF)
                self.zf = (r & 0xFFFF) == 0
                return
            if (op2 & 0xCF) == 0x4B:                          # LD rr,(nn)
                n = self.imm16()
                self.set_rr((op2 >> 4) & 3,
                            self.mem[n] | (self.mem[(n + 1) & 0xFFFF] << 8))
                return
            if (op2 & 0xCF) == 0x43:                          # LD (nn),rr
                n = self.imm16()
                v = self.get_rr((op2 >> 4) & 3)
                self.mem[n] = v & 0xFF
                self.mem[(n + 1) & 0xFFFF] = v >> 8
                return
            raise RuntimeError(f"ED {op2:02X} at {self.pc-2:04X}")

        if op == 0xDD:
            op2 = self.imm8()
            if op2 == 0x21: self.ix = self.imm16(); return    # LD IX,nn
            if op2 == 0x19:                                   # ADD IX,DE
                self.ix = (self.ix + self.de()) & 0xFFFF; return
            if op2 == 0xE5: self.push(self.ix); return
            if op2 == 0xE1: self.ix = self.pop(); return
            if op2 == 0x7E:                                   # LD A,(IX+d)
                self.a = self.mem[(self.ix + self.rel()) & 0xFFFF]; return
            if op2 == 0x77:                                   # LD (IX+d),A
                self.mem[(self.ix + self.rel()) & 0xFFFF] = self.a; return
            if op2 == 0x35:                                   # DEC (IX+d)
                p = (self.ix + self.rel()) & 0xFFFF
                self.mem[p] = (self.mem[p] - 1) & 0xFF
                self.zf = self.mem[p] == 0; return
            if op2 == 0x34:                                   # INC (IX+d)
                p = (self.ix + self.rel()) & 0xFFFF
                self.mem[p] = (self.mem[p] + 1) & 0xFF
                self.zf = self.mem[p] == 0; return
            if op2 == 0xBE:                                   # CP (IX+d)
                self.alu(7, self.mem[(self.ix + self.rel()) & 0xFFFF]); return
            if op2 == 0x36:                                   # LD (IX+d),n
                p = (self.ix + self.rel()) & 0xFFFF
                self.mem[p] = self.imm8(); return
            raise RuntimeError(f"DD {op2:02X} at {self.pc-2:04X}")

        raise RuntimeError(f"opcode {op:02X} at {self.pc-1:04X}")

    def _jr(self, take):
        d = self.rel()
        if take: self.pc = (self.pc + d) & 0xFFFF

    def _call(self, addr):
        if addr in self.call_stubs:
            self.call_stubs[addr]()
            return
        self.push(self.pc)
        self.pc = addr

    def run_frames(self, n, max_steps=20_000_000):
        """Run until n more HALTs have executed."""
        target = self.frame + n
        start = self.steps
        while self.frame < target:
            self.step()
            if self.steps - start > max_steps:
                raise RuntimeError("runaway: too many steps")
