import random
import sys

import pefile
from capstone import CS_ARCH_X86, CS_MODE_64, Cs
from capstone.x86 import X86_OP_IMM, X86_OP_REG
from unicorn import UC_ARCH_X86, UC_MODE_64, Uc
from unicorn.x86_const import UC_X86_REG_EAX, UC_X86_REG_RCX, UC_X86_REG_RSP
import llvmlite.ir as ir
import llvmlite.binding as llvm

MASK = 0xFFFFFFFF
SUBREG = {"eax": "rax", "ebx": "rbx", "ecx": "rcx", "edx": "rdx", "esi": "rsi", "edi": "rdi"}


def canon(reg):
    reg = reg.lower()
    return SUBREG.get(reg, reg)


def load_func(exe, rva):
    pe = pefile.PE(exe, fast_load=True)
    base = pe.OPTIONAL_HEADER.ImageBase
    image = pe.get_memory_mapped_image()
    return base, base + rva, image[rva:rva + 0x100]


def decode(va, code):
    md = Cs(CS_ARCH_X86, CS_MODE_64)
    md.detail = True
    out = []
    for insn in md.disasm(code, va):
        out.append(insn)
        if insn.mnemonic == "ret":
            break
    return out


def const(value):
    return ("c", value & MASK)


def pretty(node):
    kind = node[0]
    if kind == "key":
        return "key"
    if kind == "c":
        return f"0x{node[1]:X}"
    binops = {"xor": "^", "add": "+", "sub": "-", "mul": "*", "and": "&", "or": "|"}
    if kind in binops:
        return f"({pretty(node[1])} {binops[kind]} {pretty(node[2])})"
    if kind == "lshr":
        return f"({pretty(node[1])} >> {node[2]})"
    if kind == "shl":
        return f"({pretty(node[1])} << {node[2]})"
    return str(node)


def sym_exec(insns):
    regs = {"rcx": ("key",)}

    def operand(op, insn):
        if op.type == X86_OP_REG:
            return regs.get(canon(insn.reg_name(op.reg)), const(0))
        if op.type == X86_OP_IMM:
            return const(op.imm)
        return const(0)

    for insn in insns:
        mnem, ops = insn.mnemonic, insn.operands
        if mnem in ("push", "pop", "call", "nop", "ret", "endbr64", "leave"):
            continue
        if not ops or ops[0].type != X86_OP_REG:
            continue
        dst = canon(insn.reg_name(ops[0].reg))
        if dst == "rsp":
            continue
        cur = regs.get(dst, const(0))
        if mnem == "mov":
            regs[dst] = operand(ops[1], insn)
        elif mnem == "xor":
            regs[dst] = ("xor", cur, operand(ops[1], insn))
        elif mnem == "add":
            regs[dst] = ("add", cur, operand(ops[1], insn))
        elif mnem == "sub":
            regs[dst] = ("sub", cur, operand(ops[1], insn))
        elif mnem == "imul":
            lhs = operand(ops[1], insn) if len(ops) == 3 else cur
            rhs = operand(ops[2], insn) if len(ops) == 3 else operand(ops[1], insn)
            regs[dst] = ("mul", lhs, rhs)
        elif mnem == "and":
            regs[dst] = ("and", cur, operand(ops[1], insn))
        elif mnem == "or":
            regs[dst] = ("or", cur, operand(ops[1], insn))
        elif mnem == "shr":
            regs[dst] = ("lshr", cur, operand(ops[1], insn)[1])
        elif mnem == "shl":
            regs[dst] = ("shl", cur, operand(ops[1], insn)[1])
    return regs.get("rax", const(0))


def build_ir(expr):
    module = ir.Module(name="devirt")
    i32 = ir.IntType(32)
    fn = ir.Function(module, ir.FunctionType(i32, [i32]), name="transform")
    key = fn.args[0]
    key.name = "key"
    builder = ir.IRBuilder(fn.append_basic_block("entry"))

    def emit(node):
        kind = node[0]
        if kind == "key":
            return key
        if kind == "c":
            return ir.Constant(i32, node[1])
        if kind == "xor":
            return builder.xor(emit(node[1]), emit(node[2]))
        if kind == "add":
            return builder.add(emit(node[1]), emit(node[2]))
        if kind == "sub":
            return builder.sub(emit(node[1]), emit(node[2]))
        if kind == "mul":
            return builder.mul(emit(node[1]), emit(node[2]))
        if kind == "and":
            return builder.and_(emit(node[1]), emit(node[2]))
        if kind == "or":
            return builder.or_(emit(node[1]), emit(node[2]))
        if kind == "lshr":
            return builder.lshr(emit(node[1]), ir.Constant(i32, node[2]))
        if kind == "shl":
            return builder.shl(emit(node[1]), ir.Constant(i32, node[2]))
        raise ValueError(node)

    builder.ret(emit(expr))
    return str(module)


def optimize(ir_text):
    llvm.initialize_native_target()
    llvm.initialize_native_asmprinter()
    module = llvm.parse_assembly(ir_text)
    module.verify()
    machine = llvm.Target.from_default_triple().create_target_machine()
    builder = llvm.create_pass_builder(machine, llvm.create_pipeline_tuning_options(speed_level=2))
    mpm = builder.getModulePassManager()
    for name in ("add_sroa_pass", "add_instruction_combine_pass", "add_reassociate_pass",
                 "add_new_gvn_pass", "add_sccp_pass", "add_dead_code_elimination_pass",
                 "add_simplify_cfg_pass"):
        getattr(mpm, name)()
    mpm.run(module, builder)
    return module, machine


def make_jit(ir_text, machine):
    import ctypes

    module = llvm.parse_assembly(ir_text)
    engine = llvm.create_mcjit_compiler(module, machine)
    engine.finalize_object()
    engine.run_static_constructors()
    addr = engine.get_function_address("transform")
    call = ctypes.CFUNCTYPE(ctypes.c_uint32, ctypes.c_uint32)(addr)
    return engine, lambda x: call(x) & MASK


def emulate(va, code, insns, x):
    end = insns[-1].address
    buf = bytearray(code[:end - va + 1])
    i = 0
    while i < len(buf) - 1:
        if buf[i] == 0xFF and buf[i + 1] == 0x15:
            buf[i:i + 6] = b"\x90" * 6
            i += 6
        else:
            i += 1
    uc = Uc(UC_ARCH_X86, UC_MODE_64)
    uc.mem_map(va & ~0xFFF, 0x2000)
    uc.mem_write(va, bytes(buf))
    stack = 0x200000
    uc.mem_map(stack, 0x10000)
    uc.reg_write(UC_X86_REG_RSP, stack + 0x8000)
    uc.reg_write(UC_X86_REG_RCX, x)
    uc.emu_start(va, end)
    return uc.reg_read(UC_X86_REG_EAX) & MASK


def main():
    exe = sys.argv[1] if len(sys.argv) > 1 else "hello.exe"
    rva = int(sys.argv[2], 16) if len(sys.argv) > 2 else 0x10D0
    _, va, code = load_func(exe, rva)
    insns = decode(va, code)

    print(f"== {exe}: transform @ {va:#x} ({len(insns)} insns) ==\n")
    for insn in insns:
        print(f"  {insn.address:#x}  {insn.mnemonic} {insn.op_str}")

    expr = sym_exec(insns)
    print("\neax =", pretty(expr))

    ir_text = build_ir(expr)
    print("\n[unoptimized]")
    print(ir_text)

    module, machine = optimize(ir_text)
    opt_text = str(module)
    print("[optimized]")
    print(opt_text)

    engine, jit = make_jit(opt_text, machine)
    inputs = [0x1337, 0xDEADBEEF, 0, 0xFFFFFFFF] + [random.getrandbits(32) for _ in range(4)]
    ok = True
    for x in inputs:
        j, u = jit(x), emulate(va, code, insns, x)
        ok &= j == u
        print(f"key=0x{x:08X}  jit=0x{j:08X}  unicorn=0x{u:08X}  {'ok' if j == u else 'MISMATCH'}")
    print("\nverified" if ok else "\nMISMATCH")


if __name__ == "__main__":
    main()
