diff --git a/native_exec/native_function.py b/native_exec/native_function.py index 8e33ce5..92a5433 100644 --- a/native_exec/native_function.py +++ b/native_exec/native_function.py @@ -44,7 +44,6 @@ class Win32MyMap(MyMap): #return cls(-1, size, access=access) access = mmap.ACCESS_READ | mmap.ACCESS_WRITE addr = k32api.VirtualAlloc(0, size, 0x1000, 0x40) - new_map = (ctypes.c_char * size).from_address(addr) new_map.addr = addr if new_map.addr == 0: diff --git a/native_exec/simple_x64.py b/native_exec/simple_x64.py index 6fde65d..8ccf1f3 100644 --- a/native_exec/simple_x64.py +++ b/native_exec/simple_x64.py @@ -2,6 +2,8 @@ import collections import struct import sys +# TODO: fix immediat signed/unsigned assembly + class BitArray(object): def __init__(self, size, bits): self.size = size @@ -72,59 +74,133 @@ class BitArray(object): # Rules: bytes only !!!! reg_order = ['RAX', 'RCX', 'RDX', 'RBX', 'RSP', 'RBP', 'RSI', 'RDI'] - new_reg_order = ['R8', 'R9', 'R10', 'R11', 'R12', 'R13', 'R14', 'R15'] - x64_regs = reg_order + new_reg_order - -mem_access = collections.namedtuple('mem_access', ['base', 'index', 'squale', 'disp']) +mem_access = collections.namedtuple('mem_access', ['base', 'index', 'scale', 'disp']) - -def create_displacement(base=None, index=None, squale=None, disp=0): - return mem_access(base, index, squale, disp) +def create_displacement(base=None, index=None, scale=None, disp=0): + if index is not None and scale is None: + scale = 1 + return mem_access(base, index, scale, disp) + +def mem(data): + """Parse a memory access string""" + if not isinstance(data, str): + raise TypeError("mem need a string to parse") + data = data.strip() + if not (data.startswith("[") and data.endswith("]")): + raise ValueError("mem acces expect <[EXPR]>") + # A l'arrache.. j'aime pas le parsing de trucs + data = data[1:-1] + items = data.split("+") + parsed_items = {} + for item in items: + item = item.strip() + # Index * scale + if "*" in item: + if 'index' in parsed_items: + raise ValueError("Multiple index / index*scale in mem expression <{0}>".format(data)) + sub_items = item.split("*") + if len(sub_items) != 2: + raise ValueError("Invalid item <{0}> in mem access".format(item)) + index, scale = sub_items + index, scale = index.strip(), scale.strip() + if not X64.is_reg(index): + raise ValueError("Invalid index <{0}> in mem access".format(index)) + try: + scale = int(scale, 0) + except ValueError as e: + raise ValueError("Invalid scale <{0}> in mem access".format(scale)) + parsed_items['scale'] = scale + parsed_items['index'] = index + else: + # displacement / base / index alone + if X64.is_reg(item): + if not 'base' in parsed_items: + parsed_items['base'] = item + continue + # Already have base + index -> cannot avec another register in expression + if 'index' in parsed_items: + raise ValueError("Multiple index / index*scale in mem expression <{0}>".format(data)) + parsed_items['index'] = item + continue + try: + disp = int(item, 0) + except ValueError as e: + raise ValueError("Invalid base/index or displacement <{0}> in mem access".format(item)) + if 'disp' in parsed_items: + raise ValueError("Multiple displacement in mem expression <{0}>".format(data)) + parsed_items['disp'] = disp + return create_displacement(**parsed_items) class X64RegisterSelector(object): reg_opcode = {v : BitArray.from_int(size=3, x=i) for i, v in enumerate(reg_order)} new_reg_opcode = {v : BitArray.from_int(size=3, x=i) for i, v in enumerate(new_reg_order)} - + def accept_arg(self, previous, args): x = args[0] try: - return (1, self.reg_opcode[x], None) - except KeyError: + return (1, self.reg_opcode[x.upper()], None) + except (KeyError, AttributeError): pass try: - return (1, self.new_reg_opcode[x], BitArray.from_int(8, 0x41)) - except KeyError: + return (1, self.new_reg_opcode[x.upper()], BitArray.from_int(8, 0x41)) + except (KeyError, AttributeError): return (None, None, None) @classmethod def get_reg_bits(cls, name): try: - return cls.reg_opcode[name] + return cls.reg_opcode[name.upper()] except KeyError: - return cls.new_reg_opcode[name] + return cls.new_reg_opcode[name.upper()] class RawBits(BitArray): def accept_arg(self, previous, args): return (0, self, None) -class Imm64(object): + +class Immediat(object): + def __init__(self, add=0): + self.add = add + + def __add__(self, x): + return type(self)(self.add + x) + +class Imm32(Immediat): def accept_arg(self, previous, args): try: - x = int(args[0]) - return (1, BitArray.from_int(64, X64.to_little_endian(x)), None) - except TypeError: + x = int(args[0]) + self.add + except (ValueError, TypeError): + return (None, None, None) + return (1, BitArray.from_int(32, X64.to_little_endian(x, size=32)), None) + +class Imm8(Immediat): + def accept_arg(self, previous, args): + try: + x = int(args[0]) + self.add + except (ValueError, TypeError): + return (None, None) + if not -128 <= x <= 127: + return (None, None, None) + return (1, BitArray.from_int(8, X64.to_little_endian(x, size=8)), None) + +class Imm64(Immediat): + def accept_arg(self, previous, args): + try: + x = int(args[0]) + self.add + return (1, BitArray.from_int(64, X64.to_little_endian(x, size=64)), None) + except (ValueError, TypeError): return (None, None, None) class Mov_RAX_OFF64(object): def accept_arg(self, previous, args): - if args[0] != "RAX": + if RegisterRax().accept_arg(previous, args) == (None, None, None): return (None, None, None) arg2 = args[1] if not (X64.is_mem_acces(arg2) and X64.mem_access_has_only(arg2, ["disp"])): @@ -140,6 +216,22 @@ class Mov_OFF64_RAX(object): return (None, None, None) return (2, BitArray.from_int(8, 0xa3) + BitArray.from_int(64, X64.to_little_endian(arg2.disp)) , BitArray.from_int(8, 0x48)) +class RegisterRax(object): + def accept_arg(self, previous, args): + x = args[0] + if isinstance(x, str) and x.upper() == 'RAX': + return (1, BitArray(0, []), None) + return None, None, None + +class FixedRegister(object): + def __init__(self, register): + self.reg = register.upper() + + def accept_arg(self, previous, args): + x = args[0] + if isinstance(x, str) and x.upper() == self.reg: + return 1, BitArray(0, []), None + return None, None, None class ModRM(object): size = 8 @@ -172,13 +264,20 @@ class RexByte(BitArray): self.is_needed = False class X64(object): + @staticmethod def is_reg(name): - return name in x64_regs - + try: + return (name.upper() in reg_order) or X64.is_new_reg(name) + except AttributeError: # Not a string + return False + @staticmethod def is_new_reg(name): - return name in new_reg_order + try: + return name.upper() in new_reg_order + except AttributeError: # Not a string + return False @staticmethod def is_mem_acces(data): @@ -196,9 +295,12 @@ class X64(object): return True @staticmethod - def to_little_endian(i): - i = i & 0xffffffffffffffff - return struct.unpack("Q", i))[0] + def to_little_endian(i, size=64): + pack = {8: 'B', 16 : 'H', 32 : 'I', 64 : 'Q'} + s = pack[size] + mask = (1 << size) - 1 + i = i & mask + return struct.unpack("<" + s, struct.pack(">" + s, i))[0] # Sub ModRM encoding @@ -233,11 +335,12 @@ class SubModRM(object): if X64.is_new_reg(name): self.is_rex_needed = True self.rex[7] = 1 + class ModRM_REG64__REG64(SubModRM): @classmethod def match(cls, arg1, arg2): - return X64.is_reg(arg1) and X64.is_reg(arg2) + return (X64.is_reg(arg1) or X64.is_new_reg(arg1)) and (X64.is_reg(arg2) or X64.is_new_reg(arg2)) def __init__(self, arg1, arg2, reversed): super(ModRM_REG64__REG64, self).__init__() @@ -248,24 +351,11 @@ class ModRM_REG64__REG64(SubModRM): self.setup_rm_as_register(arg1) self.direction = 0 -#class ModRM_REG__DEREF_IMM(SubModRM): -# @classmethod -# def match(cls, arg1, arg2): -# return X64.is_reg(arg1) and X64.is_mem_acces(arg2) and X64.mem_access_has_only(arg2, ["disp"]) -# -# def __init__(self, arg1, arg2, reversed): -# super(ModRM_REG__DEREF_IMM, self).__init__() -# self.mod = BitArray(2, "00") -# self.setup_reg_as_register(arg1) -# self.rm = BitArray(3, "101") -# self.after = BitArray.from_int(64, X64.to_little_endian(arg2.disp)) -# self.direction = not reversed - - + class ModRM_REG__DEREF_REG(SubModRM): @classmethod def match(cls, arg1, arg2): - return X64.is_reg(arg1) and X64.is_mem_acces(arg2) and X64.mem_access_has_only(arg2, ["base"]) and arg2.base not in ["RSP", "RBP"] + return (X64.is_reg(arg1) or X64.is_new_reg(arg1)) and X64.is_mem_acces(arg2) and X64.mem_access_has_only(arg2, ["base"]) and arg2.base not in ["RSP", "RBP"] def __init__(self, arg1, arg2, reversed): super(ModRM_REG__DEREF_REG, self).__init__() @@ -277,18 +367,85 @@ class ModRM_REG__DEREF_REG(SubModRM): self.after = BitArray(0, "") self.direction = not reversed # -#class ModRM_REG__DEREF_REG_IMM(object): -# @classmethod -# def match(cls, arg1, arg2): -# return X86.is_reg(arg1) and X86.is_mem_acces(arg2) and X86.mem_access_has_only(arg2, ["base", "disp"]) -# -# def __init__(self, arg1, arg2, reversed): -# self.mod = BitArray(2, "10") -# self.reg = X86RegisterSelector.get_reg_bits(arg1) -# self.rm = X86RegisterSelector.get_reg_bits(arg2.base) -# self.after = BitArray.from_int(32, X86.to_little_endian(arg2.disp)) -# self.direction = not reversed -# +class ModRM_REG__DEREF_REG_IMM(SubModRM): + @classmethod + def match(cls, arg1, arg2): + return X64.is_reg(arg1) and X64.is_mem_acces(arg2) and X64.mem_access_has_only(arg2, ["base", "disp"]) + + def __init__(self, arg1, arg2, reversed): + super(ModRM_REG__DEREF_REG_IMM, self).__init__() + #import pdb;pdb.set_trace() + self.mod = BitArray(2, "10") + self.is_rex_needed = True + self.rex[4] = 1 + self.setup_reg_as_register(arg1) + self.setup_rm_as_register(arg2.base) + self.after = BitArray.from_int(32, X64.to_little_endian(arg2.disp, size=32)) + self.direction = not reversed + +class ModRM_REG__DEREF_BASE_INDEX(SubModRM): + """Only handle [BASE + INDEX]""" + @classmethod + def match(cls, arg1, arg2): + return X64.is_reg(arg1) and X64.is_mem_acces(arg2) and X64.mem_access_has_only(arg2, ["base", "index", "scale"]) + + def __init__(self, arg1, arg2, reversed): + super(ModRM_REG__DEREF_BASE_INDEX, self).__init__() + #import pdb;pdb.set_trace() + self.mod = BitArray(2, "00") + self.rm = BitArray(3, "100") + self.is_rex_needed = True + self.rex[4] = 1 + self.setup_reg_as_register(arg1) + self.after = self.create_sib(arg2) + self.direction = not reversed + + def create_sib(self, mem_access): + scale = {1: 0, 2 : 1, 4: 2, 8 : 3} + if mem_access.disp: + raise NotImplementedError("SIB WITH DISPLACEMENT") + if mem_access.scale not in scale: + raise ValueError("Invalid scale for mem access <{0}>".format(mem_access.scale)) + scale_bits = BitArray.from_int(2, scale[mem_access.scale]) + base_bits = X64RegisterSelector.get_reg_bits(mem_access.base) + if X64.is_new_reg(mem_access.base): + self.rex[7] = 1 + index_bits = X64RegisterSelector.get_reg_bits(mem_access.index) + if X64.is_new_reg(mem_access.index): + self.rex[6] = 1 + return scale_bits + index_bits + base_bits + + +class Slash(object): + "No idea for the name: represent the modRM for single args + encoding in reg (/7 in cmp in man intel)" + + def __init__(self, reg): + "reg = 7 for /7" + self.mod = None + self.reg = BitArray.from_int(3, reg) + + self.rm = None + + def accept_arg(self, previous, args): + x = args[0] + ok, bits, rex = X64RegisterSelector().accept_arg(None, [x]) + if ok is not None: + self.mod = BitArray(2, "11") + self.rm = bits + return 1, self.mod + self.reg + self.rm, rex + # TODO: register ! + if X64.mem_access_has_only(x, ["base"]) and x.base not in ['ESP', 'EBP']: + self.mod = BitArray(2, "00") + ok, bits = X86RegisterSelector().accept_arg(None, [x.base]) + self.rm = bits + return 1, self.mod + self.reg + self.rm, rex + # TODO: Other + if X64.mem_access_has_only(x, ["base", "disp"]): + self.mod = BitArray(2, "10") + ok, bits, rex = X64RegisterSelector().accept_arg(None, [x.base]) + self.rm = bits + return 1, self.mod + self.reg + self.rm + BitArray.from_int(32, X64.to_little_endian(x.disp, size=32)), rex + return None, None class Instruction(object): encoding = [] @@ -315,12 +472,29 @@ class Instruction(object): if any(full_rex.array): self.value = full_rex + self.value return - raise ValueError("Cannot encode :(") + raise ValueError("Cannot encode <{0} {1}>:(".format(type(self).__name__, initial_args)) + + def get_code(self): + return self.value.dump() + +class DelayedJump(object): + def __init__(self, type, label): + self.type = type + self.label = label + +class JmpType(Instruction): + def __new__(cls, *initial_args): + if len(initial_args) == 1: + arg = initial_args[0] + if isinstance(arg, str) and arg[0] == ":": + return DelayedJump(cls, arg) + return super(JmpType, cls).__new__(cls, *initial_args) + # # class Push(Instruction): - encoding = [(RawBits.from_int(5, 0x50 >> 3), X64RegisterSelector()),] -# (RawBits.from_int(8, 0x68), Imm32())] + encoding = [(RawBits.from_int(5, 0x50 >> 3), X64RegisterSelector()), + (RawBits.from_int(8, 0x68), Imm32())] class Pop(Instruction): encoding = [(RawBits.from_int(5, 0x58 >> 3), X64RegisterSelector())] @@ -334,25 +508,306 @@ class Ret(Instruction): class Int3(Instruction): encoding = [(RawBits.from_int(8, 0xcc),)] +class Dec(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0xff), Slash(1))] + +class Inc(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0xff), Slash(0))] + +class Add(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0x05), RegisterRax(), Imm32()), + (RawBits.from_int(8, 0x81), Slash(0), Imm32()), + (RawBits.from_int(8, 0x01), ModRM(ModRM_REG64__REG64, ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_BASE_INDEX)),] + +class Sub(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0x2D), RegisterRax(), Imm32()), + (RawBits.from_int(8, 0x81), Slash(5), Imm32())] + +class Out(Instruction): + encoding = [(RawBits.from_int(8, 0xee), FixedRegister('DX'), FixedRegister('AL')), + (RawBits.from_int(16, 0x66ef), FixedRegister('DX'), FixedRegister('AX')), # Fuck-it hardcoded prefix for now + (RawBits.from_int(8, 0xef), FixedRegister('DX'), FixedRegister('EAX'))] + +class In(Instruction): + encoding = [(RawBits.from_int(8, 0xec), FixedRegister('AL'), FixedRegister('DX')), + (RawBits.from_int(16, 0x66ed), FixedRegister('AX'), FixedRegister('DX')), # Fuck-it hardcoded prefix for now + (RawBits.from_int(8, 0xed), FixedRegister('EAX'), FixedRegister('DX'))] + +class JmpImm8(Immediat): + def __init__(self, sub): + self.sub = sub + def accept_arg(self, previous, args): + try: + x = int(args[0]) + except (ValueError, TypeError): + return (None, None, None) + if not (-128 + self.sub) <= x <= 127: + return (None, None, None) + x -= self.sub + return (1, BitArray.from_int(8, X64.to_little_endian(x, size=8)), None) + +class JmpImm32(Immediat): + def __init__(self, sub): + self.sub = sub + def accept_arg(self, previous, args): + try: + x = int(args[0]) + except (ValueError, TypeError): + return (None, None, None) + #if not (-128 + self.ADD) <= x <= 127: + # return (None, None) + x -= self.sub + return (1, BitArray.from_int(32, X64.to_little_endian(x, size=32)), None) + +class Jmp(JmpType): + encoding = [(RawBits.from_int(8, 0xeb), JmpImm8(2)), + (RawBits.from_int(8, 0xe9), JmpImm32(5)), + (RawBits.from_int(13, 0xffe0 >> 3), X64RegisterSelector())] + +class Jz(JmpType): + encoding = [(RawBits.from_int(8, 0x74), JmpImm8(2)), + (RawBits.from_int(16, 0x0f84), JmpImm32(6))] + +class Jnz(JmpType): + encoding = [(RawBits.from_int(8, 0x75), JmpImm8(2)), + (RawBits.from_int(16, 0x0f85), JmpImm32(6))] + +class Jb(JmpType): + encoding = [(RawBits.from_int(8, 0x72), JmpImm8(2)), + (RawBits.from_int(16, 0x0f82), JmpImm32(6))] + +class Jbe(JmpType): + encoding = [(RawBits.from_int(8, 0x76), JmpImm8(2)), + (RawBits.from_int(16, 0x0f86), JmpImm32(6))] + +class Jnb(JmpType): + encoding = [(RawBits.from_int(8, 0x73), JmpImm8(2)), + (RawBits.from_int(16, 0x0f83), JmpImm32(6))] + class Mov(Instruction): default_32_bits = True - encoding = [(RawBits.from_int(8, 0x89), ModRM(ModRM_REG64__REG64, ModRM_REG__DEREF_REG)), (RawBits.from_int(5, 0xb8 >> 3), X64RegisterSelector(), Imm64()), + encoding = [(RawBits.from_int(8, 0x89), ModRM(ModRM_REG64__REG64, ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_BASE_INDEX)), + (RawBits.from_int(5, 0xb8 >> 3), X64RegisterSelector(), Imm64()), (Mov_RAX_OFF64(),), (Mov_OFF64_RAX(),)] - + +class Cmp(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0x3d), RegisterRax(), Imm32()), + (RawBits.from_int(8, 0x81), Slash(7), Imm32()), + (RawBits.from_int(8, 0x3b), ModRM(ModRM_REG64__REG64, ModRM_REG__DEREF_REG)),] + +class Xor(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0x31), ModRM(ModRM_REG64__REG64))] + +class Nop(Instruction): + encoding = [(RawBits.from_int(8, 0x90),)] + +class Retf(Instruction): + default_32_bits = True + encoding = [(RawBits.from_int(8, 0xcb),)] + +class Retf32(Instruction): + encoding = [(RawBits.from_int(8, 0xcb),)] + +class _NopArtifact(Nop): + pass + +def JmpAt(addr): + code = MultipleInstr() + code += Mov('RAX', addr) + code += Jmp('RAX') + return code + +class Label(object): + def __init__(self, name): + self.name = name + class MultipleInstr(object): - - def __init__(self, instrs=()): - self.instrs = list(instrs) - - def __iadd__(self, value): - if type(value) == MultipleInstr: - self.instrs.extend(value.instrs) - return self - self.instrs.append(value) - return self + JUMP_SIZE = 6 + def __init__(self, init_instrs=()): + self.instrs = {} + self.labels = {} + self.expected_labels = {} + # List of all labeled jump already resolved + # Will be used for 'relocation' + self.computed_jump = [] + self.size = 0 + for i in init_instrs: + self += i def get_code(self): - if sys.version_info.major == 3: - return b"".join([x.value.dump() for x in self.instrs]) - return "".join([str(x.value.dump()) for x in self.instrs]) + if self.expected_labels: + raise ValueError("Unresolved labels: {self.expected_labels}".format(self=self)) + return "".join([str(x[1].get_code()) for x in sorted(self.instrs.items())]) + + def add_instruction(self, instruction): + if isinstance(instruction, Label): + return self.add_label(instruction) + # Change DelayedJump to LabeledJump ? + if isinstance(instruction, DelayedJump): + return self.add_delayed_jump(instruction) + if isinstance(instruction, Instruction): + self.instrs[self.size] = instruction + self.size += len(instruction.get_code()) + return + raise ValueError("Don't know what to do with {0} of type {1}".format(instruction, type(instruction))) + + def add_label(self, label): + if label.name not in self.expected_labels: + # Label that have no jump before definition + # Just registed the address of the label + self.labels[label.name] = self.size + return + # Label with jmp before definition + # Lot of stuff todo: + # Find all delayed jump that refer to this jump + # Replace them with real jump + # If size of jump < JUMP_SIZE: relocate everything we can + # Update expected_labels + for jump_to_label in self.expected_labels[label.name]: + if jump_to_label.offset in self.instrs: + raise ValueError("WTF REPLACE EXISTING INSTR...") + distance = self.size - jump_to_label.offset + real_jump = jump_to_label.type(distance) + self.instrs[jump_to_label.offset] = real_jump + self.computed_jump.append((jump_to_label.offset, self.size)) + for i in range(self.JUMP_SIZE - len(real_jump.get_code())): + self.instrs[jump_to_label.offset + len(real_jump.get_code()) + i] = _NopArtifact() + del self.expected_labels[label.name] + self.labels[label.name] = self.size + if not self.expected_labels: + # No more un-resolved label (for now): time to reduce the shellcode + self._reduce_shellcode() + + def add_delayed_jump(self, jump): + dst = jump.label + if dst in self.labels: + # Jump to already defined labels + # Nothing fancy: get offset of label and jump to it ! + distance = self.size - self.labels[dst] + jump_instruction = jump.type(-distance) + self.computed_jump.append((self.size, self.labels[dst])) + return self.add_instruction(jump_instruction) + # Jump to undefined label + # Add label to expected ones + # Add jump info -> offset of jump | type + # Reserve space for call ! + jump.offset = self.size + self.expected_labels.setdefault(dst, []).append(jump) + self.size += self.JUMP_SIZE + return + + def _reduce_shellcode(self): + to_remove = [offset for offset,instr in self.instrs.items() if type(instr) == _NopArtifact] + while to_remove: + self._remove_nop_artifact(to_remove[0]) + # _remove_nop_artifact will change the offsets of the nop + # Need to refresh these offset + to_remove = [offset for offset,instr in self.instrs.items() if type(instr) == _NopArtifact] + + def _remove_nop_artifact(self, offset): + # Remove a NOP from the shellcode + for src, dst in self.computed_jump: + # Reduce size of Jump over the nop (both sens) + if src < offset < dst or dst < offset < src: + old_jmp = self.instrs[src] + old_jump_size = len(old_jmp.get_code()) + if src < offset < dst: + new_jmp = type(old_jmp)(dst - src - 1) + else: + new_jmp = type(old_jmp)(dst - src + 1) + new_jmp_size = len(new_jmp.get_code()) + if new_jmp_size > old_jump_size: + raise ValueError("Wtf jump of smaller size of bigger.. ABORT") + self.instrs[src] = new_jmp + # Add other _NopArtifact if jump instruction size is reduced + for i in range(old_jump_size - new_jmp_size): + self.instrs[src + new_jmp_size + i] = _NopArtifact() + + # dec offset of all Label after the NOP + for name, labeloffset in self.labels.items(): + if labeloffset > offset: + self.labels[name] = labeloffset - 1 + + # dec offset of all instr after the NOP + new_instr = {} + for instroffset, instr in self.instrs.items(): + if instroffset == offset: + continue + if instroffset > offset: + instroffset -= 1 + new_instr[instroffset] = instr + self.instrs = new_instr + # Update all computed jump + new_computed_jump = [] + for src, dst in self.computed_jump: + if src > offset: + src -= 1 + if dst > offset: + dst -= 1 + new_computed_jump.append((src, dst)) + self.computed_jump = new_computed_jump + # dec size of the shellcode + self.size -= 1 + + def merge_shellcode(self, other): + for offset, instr in sorted(other.instrs.items()): + self.add_instruction(instr) + + def __iadd__(self, other): + if isinstance(other, MultipleInstr): + self.merge_shellcode(other) + else: + self.add_instruction(other) + return self + + +# import windows.native_exec.simple_x64 as x64 +try: + import midap + import idc + in_IDA = True +except ImportError: + in_IDA = False + + +def test_code(): + s = MultipleInstr() + s += Mov('r8', 'r14') + s += Label(':SUCE') + s += Jnz(':END') + s += Add('r14', 0x12345678) + s += Dec('r9') + s += Dec('rax') + s += Jnz(':END') + s += Mov('r8', 'rdx') + s += Jnz(':END') + s += Mov('r8', 'rdx') + s += Jnz(':SUCE') + s += Mov('r9', 'r10') + s += Label(':END') + s += Ret() + return s + + + +if in_IDA: + def reset(): + idc.MakeUnknown(idc.MinEA(), 0x1000, 0) + for i in range(0x1000): + idc.PatchByte(idc.MinEA() + i, 0) + + s = test_code() + + def tst(): + reset() + midap.here(idc.MinEA()).write(s.get_code()) + idc.MakeFunction(idc.MinEA()) + + #tst() diff --git a/native_exec/simple_x86.py b/native_exec/simple_x86.py index 607d736..5e2f31c 100644 --- a/native_exec/simple_x86.py +++ b/native_exec/simple_x86.py @@ -63,13 +63,65 @@ class BitArray(object): # Rules: bytes only !!!! -mem_access = collections.namedtuple('mem_access', ['base', 'index', 'squale', 'disp']) +mem_access = collections.namedtuple('mem_access', ['base', 'index', 'scale', 'disp']) x86_regs = ['EAX', 'ECX', 'EDX', 'EBX', 'ESP', 'EBP', 'ESI', 'EDI'] -def create_displacement(base=None, index=None, squale=None, disp=0): - return mem_access(base, index, squale, disp) - - +def create_displacement(base=None, index=None, scale=None, disp=0): + if index is not None and scale is None: + scale = 1 + return mem_access(base, index, scale, disp) + +def mem(data): + """Parse a memory access string""" + if not isinstance(data, str): + raise TypeError("mem need a string to parse") + data = data.strip() + if not (data.startswith("[") and data.endswith("]")): + raise ValueError("mem acces expect <[EXPR]>") + # A l'arrache.. j'aime pas le parsing de trucs + data = data[1:-1] + items = data.split("+") + parsed_items = {} + for item in items: + item = item.strip() + # Index * scale + if "*" in item: + if 'index' in parsed_items: + raise ValueError("Multiple index / index*scale in mem expression <{0}>".format(data)) + sub_items = item.split("*") + if len(sub_items) != 2: + raise ValueError("Invalid item <{0}> in mem access".format(item)) + index, scale = sub_items + index, scale = index.strip(), scale.strip() + if not X86.is_reg(index): + raise ValueError("Invalid index <{0}> in mem access".format(index)) + try: + scale = int(scale, 0) + except ValueError as e: + raise ValueError("Invalid scale <{0}> in mem access".format(scale)) + parsed_items['scale'] = scale + parsed_items['index'] = index + else: + # displacement / base / index alone + if X86.is_reg(item): + if not 'base' in parsed_items: + parsed_items['base'] = item + continue + # Already have base + index -> cannot avec another register in expression + if 'index' in parsed_items: + raise ValueError("Multiple index / index*scale in mem expression <{0}>".format(data)) + parsed_items['index'] = item + continue + try: + disp = int(item, 0) + except ValueError as e: + raise ValueError("Invalid base/index or displacement <{0}> in mem access".format(item)) + if 'disp' in parsed_items: + raise ValueError("Multiple displacement in mem expression <{0}>".format(data)) + parsed_items['disp'] = disp + return create_displacement(**parsed_items) + + class X86RegisterSelector(object): size = 3 # bits reg_order = ['EAX', 'ECX', 'EDX', 'EBX', 'ESP', 'EBP', 'ESI', 'EDI'] @@ -78,30 +130,67 @@ class X86RegisterSelector(object): def accept_arg(self, previous, args): x = args[0] try: - return (1, self.reg_opcode[x]) - except KeyError: + return (1, self.reg_opcode[x.upper()]) + except (KeyError, AttributeError): return (None, None) @classmethod def get_reg_bits(cls, name): - return cls.reg_opcode[name] + return cls.reg_opcode[name.upper()] + +class RegisterEax(object): + def accept_arg(self, previous, args): + x = args[0] + if isinstance(x, str) and x.upper() == 'EAX': + return (1, BitArray(0, [])) + return None, None + +class FixedRegister(object): + def __init__(self, register): + self.reg = register.upper() + + def accept_arg(self, previous, args): + x = args[0] + if isinstance(x, str) and x.upper() == self.reg: + return (1, BitArray(0, [])) + return None, None + class RawBits(BitArray): def accept_arg(self, previous, args): return (0, self) -class Imm32(object): +class Immediat(object): + def __init__(self, add=0): + self.add = add + + def __add__(self, x): + return type(self)(self.add + x) + +class Imm32(Immediat): def accept_arg(self, previous, args): try: - x = int(args[0]) - except TypeError: + x = int(args[0]) + self.add + except (ValueError, TypeError): return (None, None) - return (1, BitArray.from_int(32, X86.to_little_endian(x))) + return (1, BitArray.from_int(32, X86.to_little_endian(x, size=32))) + +class Imm8(Immediat): + def accept_arg(self, previous, args): + try: + x = int(args[0]) + self.add + except (ValueError, TypeError): + return (None, None) + if not -128 <= x <= 127: + return (None, None) + return (1, BitArray.from_int(8, X86.to_little_endian(x, size=8))) class ModRM(object): size = 8 - def __init__(self, *sub_modrm): + def __init__(self, sub_modrm, accept_reverse=True, has_direction_bit=True): + self.accept_reverse = accept_reverse + self.has_direction_bit = has_direction_bit self.sub = sub_modrm def accept_arg(self, previous, args): @@ -110,21 +199,27 @@ class ModRM(object): arg1 = args[0] arg2 = args[1] for sub in self.sub: + # Problem in reverse sens -> need to fix it #import pdb;pdb.set_trace() if sub.match(arg1, arg2): d = sub(arg1, arg2, 0) - previous[0][-2] = d.direction + if self.has_direction_bit: + previous[0][-2] = d.direction return (2, d.mod + d.reg + d.rm + d.after) - elif sub.match(arg2, arg1): + elif self.accept_reverse and sub.match(arg2, arg1): d = sub(arg2, arg1, 1) - previous[0][-2] = d.direction + if self.has_direction_bit: + previous[0][-2] = d.direction return (2, d.mod + d.reg + d.rm + d.after) return (None, None) class X86(object): @staticmethod def is_reg(name): - return name in x86_regs + try: + return name.upper() in x86_regs + except AttributeError: # Not a string + return False @staticmethod def is_mem_acces(data): @@ -143,9 +238,12 @@ class X86(object): return True @staticmethod - def to_little_endian(i): - i = i & 0xffffffff - return struct.unpack("I", i))[0] + def to_little_endian(i, size=32): + pack = {8: 'B', 16 : 'H', 32 : 'I'} + s = pack[size] + mask = (1 << size) - 1 + i = i & mask + return struct.unpack("<" + s, struct.pack(">" + s, i))[0] class ModRM_REG__REG(object): @@ -184,21 +282,34 @@ class ModRM_REG__DEREF_REG_IMM(object): self.after = BitArray.from_int(32, X86.to_little_endian(arg2.disp)) self.direction = not reversed +def sib_from_mem_access(mem_access): + scale = {1: 0, 2 : 1, 4: 2, 8 : 3} + if mem_access.scale is None and mem_access.index is None: + return BitArray.from_int(2, 0) + BitArray.from_int(3, 0b100) + X86RegisterSelector.get_reg_bits(mem_access.base) + if mem_access.scale not in scale: + raise ValueError("Invalid scale for mem access <{0}>".format(mem_access.scale)) + return BitArray.from_int(2, scale[mem_access.scale]) + X86RegisterSelector.get_reg_bits(mem_access.index) + X86RegisterSelector.get_reg_bits(mem_access.base) + + class ModRM_REG__DEREF_SIB(object): # Only handle reg, [esp+x] now :( @classmethod def match(cls, arg1, arg2): - return X86.is_reg(arg1) and X86.is_mem_acces(arg2) and X86.mem_access_has_only(arg2, ["base", "disp"]) and arg2.base == "ESP" + return X86.is_reg(arg1) and X86.is_mem_acces(arg2) def __init__(self, arg1, arg2, reversed): - self.mod = BitArray(2, "10") + if not arg2.disp: + self.mod = BitArray(2, "00") + else: + self.mod = BitArray(2, "10") self.reg = X86RegisterSelector.get_reg_bits(arg1) self.rm = BitArray(3, "100") - # Todo -> def sib_from_displacement - sib = BitArray(8, "00100100") - - self.after = sib + BitArray.from_int(32, X86.to_little_endian(arg2.disp)) + sib = sib_from_mem_access(arg2) + if arg2.disp: + self.after = sib + BitArray.from_int(32, X86.to_little_endian(arg2.disp)) + else: + self.after = sib self.direction = not reversed #class ModRM_REG_IMM(object): @@ -223,8 +334,38 @@ class ModRM_REG__DEREF_IMM(object): self.rm = BitArray(3, "101") self.after = BitArray.from_int(32, X86.to_little_endian(arg2.disp)) self.direction = not reversed + +class Slash(object): + "No idea for the name: represent the modRM for single args + encoding in reg (/7 in cmp in man intel)" - + def __init__(self, reg): + "reg = 7 for /7" + self.mod = None + self.reg = BitArray.from_int(3, reg) + self.rm = None + + def accept_arg(self, previous, args): + x = args[0] + ok, bits = X86RegisterSelector().accept_arg(None, [x]) + if ok is not None: + self.mod = BitArray(2, "11") + self.rm = bits + return 1, self.mod + self.reg + self.rm + if X86.mem_access_has_only(x, ["base"]) and x.base not in ['ESP', 'EBP']: + self.mod = BitArray(2, "00") + ok, bits = X86RegisterSelector().accept_arg(None, [x.base]) + self.rm = bits + return 1, self.mod + self.reg + self.rm + # TODO: Other + + if X86.mem_access_has_only(x, ["base", "disp"]): + self.mod = BitArray(2, "10") + ok, bits = X86RegisterSelector().accept_arg(None, [x.base]) + self.rm = bits + return 1, self.mod + self.reg + self.rm + BitArray.from_int(32, X86.to_little_endian(x.disp)) + return None, None + + class Instruction(object): encoding = [] @@ -243,9 +384,25 @@ class Instruction(object): continue self.value = sum(res, BitArray(0, "")) return - raise ValueError("Cannot encode :(") - - + raise ValueError("Cannot encode <{0} {1}>:(".format(type(self).__name__, initial_args)) + + def get_code(self): + return self.value.dump() + +class DelayedJump(object): + def __init__(self, type, label): + self.type = type + self.label = label + +class JmpType(Instruction): + def __new__(cls, *initial_args): + if len(initial_args) == 1: + arg = initial_args[0] + if isinstance(arg, str) and arg[0] == ":": + return DelayedJump(cls, arg) + return super(JmpType, cls).__new__(cls, *initial_args) + + class Push(Instruction): encoding = [(RawBits.from_int(5, 0x50 >> 3), X86RegisterSelector()), (RawBits.from_int(8, 0x68), Imm32())] @@ -253,34 +410,305 @@ class Push(Instruction): class Pop(Instruction): encoding = [(RawBits.from_int(5, 0x58 >> 3), X86RegisterSelector())] +class Dec(Instruction): + encoding = [(RawBits.from_int(5, 0x48 >> 3), X86RegisterSelector())] + +class Inc(Instruction): + encoding = [(RawBits.from_int(5, 0x40 >> 3), X86RegisterSelector())] + +class Add(Instruction): + encoding = [(RawBits.from_int(8, 0x05), RegisterEax(), Imm32()), + (RawBits.from_int(8, 0x81), Slash(0), Imm32()), + (RawBits.from_int(8, 0x01), ModRM([ModRM_REG__REG, ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_IMM, ModRM_REG__DEREF_SIB])),] + +class Sub(Instruction): + encoding = [(RawBits.from_int(8, 0x2D), RegisterEax(), Imm32()), + (RawBits.from_int(8, 0x81), Slash(5), Imm32())] + class Mov(Instruction): - encoding = [(RawBits.from_int(8, 0x89), ModRM(ModRM_REG__REG, ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_IMM, ModRM_REG__DEREF_SIB)), + encoding = [(RawBits.from_int(8, 0x89), ModRM([ModRM_REG__REG, ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_IMM, ModRM_REG__DEREF_SIB])), (RawBits.from_int(5, 0xb8 >> 3), X86RegisterSelector(), Imm32())] + +class Lea(Instruction): + encoding = [(RawBits.from_int(8, 0x8d), ModRM([ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_IMM, ModRM_REG__DEREF_SIB], accept_reverse=False, has_direction_bit=False))] class Call(Instruction): encoding = [(RawBits.from_int(13, 0xffd0 >> 3), X86RegisterSelector())] +class Cmp(Instruction): + encoding = [(RawBits.from_int(8, 0x3d), RegisterEax(), Imm32()), + (RawBits.from_int(8, 0x81), Slash(7), Imm32()), + (RawBits.from_int(8, 0x3b), ModRM([ModRM_REG__REG, ModRM_REG__DEREF_REG, ModRM_REG__DEREF_REG_IMM, ModRM_REG__DEREF_IMM, ModRM_REG__DEREF_SIB])), ] + +class Out(Instruction): + encoding = [(RawBits.from_int(8, 0xee), FixedRegister('DX'), FixedRegister('AL')), + (RawBits.from_int(16, 0x66ef), FixedRegister('DX'), FixedRegister('AX')), # Fuck-it hardcoded prefix for now + (RawBits.from_int(8, 0xef), FixedRegister('DX'), FixedRegister('EAX'))] + +class In(Instruction): + encoding = [(RawBits.from_int(8, 0xec), FixedRegister('AL'), FixedRegister('DX')), + (RawBits.from_int(16, 0x66ed), FixedRegister('AX'), FixedRegister('DX')), # Fuck-it hardcoded prefix for now + (RawBits.from_int(8, 0xed), FixedRegister('EAX'), FixedRegister('DX'))] + + +class JmpImm8(Immediat): + def __init__(self, sub): + self.sub = sub + def accept_arg(self, previous, args): + try: + x = int(args[0]) + except (ValueError, TypeError): + return (None, None) + if not (-128 + self.sub) <= x <= 127: + return (None, None) + x -= self.sub + return (1, BitArray.from_int(8, X86.to_little_endian(x, size=8))) + +class JmpImm32(Immediat): + def __init__(self, sub): + self.sub = sub + def accept_arg(self, previous, args): + try: + x = int(args[0]) + except (ValueError, TypeError): + return (None, None) + #if not (-128 + self.ADD) <= x <= 127: + # return (None, None) + x -= self.sub + return (1, BitArray.from_int(32, X86.to_little_endian(x, size=32))) + +class Jmp(JmpType): + encoding = [(RawBits.from_int(8, 0xeb), JmpImm8(2)), + (RawBits.from_int(8, 0xe9), JmpImm32(5))] + +class Jz(JmpType): + encoding = [(RawBits.from_int(8, 0x74), JmpImm8(2)), + (RawBits.from_int(16, 0x0f84), JmpImm32(6))] + +class Jnz(JmpType): + encoding = [(RawBits.from_int(8, 0x75), JmpImm8(2)), + (RawBits.from_int(16, 0x0f85), JmpImm32(6))] + +class Xor(Instruction): + encoding = [(RawBits.from_int(8, 0x31), ModRM([ModRM_REG__REG]))] + class Ret(Instruction): encoding = [(RawBits.from_int(8, 0xc3),)] +class Nop(Instruction): + encoding = [(RawBits.from_int(8, 0x90),)] + +class Retf(Instruction): + encoding = [(RawBits.from_int(8, 0xcb),)] + +class _NopArtifact(Nop): + pass + class Int3(Instruction): encoding = [(RawBits.from_int(8, 0xcc),)] +class Label(object): + def __init__(self, name): + self.name = name +def JmpAt(addr): + code = MultipleInstr() + code += Push(addr) + code += Ret() + return code + class MultipleInstr(object): - - def __init__(self): - self.instrs = [] - - def __iadd__(self, value): - if type(value) == MultipleInstr: - self.instrs.extend(value.instrs) - return self - self.instrs.append(value) - return self + JUMP_SIZE = 6 + def __init__(self, init_instrs=()): + self.instrs = {} + self.labels = {} + self.expected_labels = {} + # List of all labeled jump already resolved + # Will be used for 'relocation' + self.computed_jump = [] + self.size = 0 + for i in init_instrs: + self += i def get_code(self): - if sys.version_info.major == 3: - return b"".join([x.value.dump() for x in self.instrs]) - return "".join([str(x.value.dump()) for x in self.instrs]) + if self.expected_labels: + raise ValueError("Unresolved labels: {self.expected_labels}".format(self=self)) + return "".join([str(x[1].get_code()) for x in sorted(self.instrs.items())]) + + def add_instruction(self, instruction): + if isinstance(instruction, Label): + return self.add_label(instruction) + # Change DelayedJump to LabeledJump ? + if isinstance(instruction, DelayedJump): + return self.add_delayed_jump(instruction) + if isinstance(instruction, Instruction): + self.instrs[self.size] = instruction + self.size += len(instruction.get_code()) + return + raise ValueError("Don't know what to do with {0} of type {1}".format(instruction, type(instruction))) + + def add_label(self, label): + if label.name not in self.expected_labels: + # Label that have no jump before definition + # Just registed the address of the label + self.labels[label.name] = self.size + return + # Label with jmp before definition + # Lot of stuff todo: + # Find all delayed jump that refer to this jump + # Replace them with real jump + # If size of jump < JUMP_SIZE: relocate everything we can + # Update expected_labels + for jump_to_label in self.expected_labels[label.name]: + if jump_to_label.offset in self.instrs: + raise ValueError("WTF REPLACE EXISTING INSTR...") + distance = self.size - jump_to_label.offset + real_jump = jump_to_label.type(distance) + self.instrs[jump_to_label.offset] = real_jump + self.computed_jump.append((jump_to_label.offset, self.size)) + for i in range(self.JUMP_SIZE - len(real_jump.get_code())): + self.instrs[jump_to_label.offset + len(real_jump.get_code()) + i] = _NopArtifact() + del self.expected_labels[label.name] + self.labels[label.name] = self.size + if not self.expected_labels: + # No more un-resolved label (for now): time to reduce the shellcode + self._reduce_shellcode() + + def add_delayed_jump(self, jump): + dst = jump.label + if dst in self.labels: + # Jump to already defined labels + # Nothing fancy: get offset of label and jump to it ! + distance = self.size - self.labels[dst] + jump_instruction = jump.type(-distance) + self.computed_jump.append((self.size, self.labels[dst])) + return self.add_instruction(jump_instruction) + # Jump to undefined label + # Add label to expected ones + # Add jump info -> offset of jump | type + # Reserve space for call ! + jump.offset = self.size + self.expected_labels.setdefault(dst, []).append(jump) + self.size += self.JUMP_SIZE + return + + def _reduce_shellcode(self): + to_remove = [offset for offset,instr in self.instrs.items() if type(instr) == _NopArtifact] + while to_remove: + self._remove_nop_artifact(to_remove[0]) + # _remove_nop_artifact will change the offsets of the nop + # Need to refresh these offset + to_remove = [offset for offset,instr in self.instrs.items() if type(instr) == _NopArtifact] + + def _remove_nop_artifact(self, offset): + # Remove a NOP from the shellcode + for src, dst in self.computed_jump: + # Reduce size of Jump over the nop (both sens) + if src < offset < dst or dst < offset < src: + old_jmp = self.instrs[src] + old_jump_size = len(old_jmp.get_code()) + if src < offset < dst: + new_jmp = type(old_jmp)(dst - src - 1) + else: + new_jmp = type(old_jmp)(dst - src + 1) + new_jmp_size = len(new_jmp.get_code()) + if new_jmp_size > old_jump_size: + raise ValueError("Wtf jump of smaller size of bigger.. ABORT") + self.instrs[src] = new_jmp + # Add other _NopArtifact if jump instruction size is reduced + for i in range(old_jump_size - new_jmp_size): + self.instrs[src + new_jmp_size + i] = _NopArtifact() + + # dec offset of all Label after the NOP + for name, labeloffset in self.labels.items(): + if labeloffset > offset: + self.labels[name] = labeloffset - 1 + + # dec offset of all instr after the NOP + new_instr = {} + for instroffset, instr in self.instrs.items(): + if instroffset == offset: + continue + if instroffset > offset: + instroffset -= 1 + new_instr[instroffset] = instr + self.instrs = new_instr + # Update all computed jump + new_computed_jump = [] + for src, dst in self.computed_jump: + if src > offset: + src -= 1 + if dst > offset: + dst -= 1 + new_computed_jump.append((src, dst)) + self.computed_jump = new_computed_jump + # dec size of the shellcode + self.size -= 1 + + def merge_shellcode(self, other): + for offset, instr in sorted(other.instrs.items()): + self.add_instruction(instr) + + def __iadd__(self, other): + if isinstance(other, MultipleInstr): + self.merge_shellcode(other) + else: + self.add_instruction(other) + return self + +# IDA : import windows.native_exec.simple_x86 as x86 +# IDA testing + +try: + import midap + import idc + in_IDA = True +except ImportError: + in_IDA = False + +#def test_code(): +# s = MultipleInstr() +# s += Mov('EAX', 'EAX') +# s += Mov('EAX', 'EAX') +# s += Jnz(":SUCE") +# s += Mov('EAX', 'EAX') +# s += Cmp("Eax", "ESI") +# s += Jnz(":SUCE") +# s += Mov("ECX", "ECX") +# s += Label(":SUCE") +# s += Jnz(":LOL") +# s += Jnz(":BITE") +# s += Mov("EDX", "EDX") +# s += Label(":LOL") +# s += Mov('EDI', 'EDI') +# s += Label(":BITE") +# s += Mov('EDI', 'EDI') +# s += Jnz(":SUCE") +# s += Push("ECX") +# s += Pop("EAX") +# s += Ret() +# return s + +def test_code(): + s = MultipleInstr() + s += Mov("Eax", "ESI") + s += Inc("Ecx") + s += Dec("edi") + s += Ret() + return s + + + +if in_IDA: + def reset(): + idc.MakeUnknown(idc.MinEA(), 0x1000, 0) + for i in range(0x1000): + idc.PatchByte(idc.MinEA() + i, 0) + + s = test_code() + + def tst(): + reset() + midap.here(idc.MinEA()).write(s.get_code()) + idc.MakeFunction(idc.MinEA()) \ No newline at end of file