import gdb
import os


GSP_CONFIG_OFFSETS = {0x4018, 0x401c, 0x4020, 0x1b400}
GSP_ONLY = os.environ.get("IPU_TRACE_GSP_ONLY") == "1"


class SyncBreakpoint(gdb.Breakpoint):
    def __init__(self, address, label, arguments, image_base=None):
        super().__init__("*0x%x" % address, internal=False)
        self.label = label
        self.arguments = arguments
        self.image_base = image_base

    def stop(self):
        values = []
        for name, register in self.arguments:
            value = int(gdb.parse_and_eval("$" + register))
            values.append("%s=%#x" % (name, value))
        if self.image_base is not None:
            frame = gdb.newest_frame().older()
            callers = []
            while frame is not None and len(callers) != 3:
                callers.append("%#x" % (frame.pc() - self.image_base))
                frame = frame.older()
            values.append("callers=" + ",".join(callers))
        gdb.write("%s %s\n" % (self.label, " ".join(values)))
        return False


class ResolveDeviceSetMark(gdb.Breakpoint):
    def __init__(self, address):
        super().__init__("*0x%x" % address, internal=False, temporary=True)

    def stop(self):
        self.enabled = False
        device = int(gdb.parse_and_eval("$rdi"))
        inferior = gdb.selected_inferior()
        pointer_size = gdb.lookup_type("void").pointer().sizeof
        vtable = int.from_bytes(inferior.read_memory(device, pointer_size),
                                byteorder="little")
        target = int.from_bytes(inferior.read_memory(vtable + 0x2B8,
                                                     pointer_size),
                                byteorder="little")
        gdb.write("device-set-mark target=%#x library=%s\n" %
                  (target, gdb.solib_name(target)))
        SyncBreakpoint(target, "device-set-mark",
                       [("this", "rdi"), ("proxy", "rsi"),
                        ("mark", "rdx"), ("ipu", "rcx")])
        return False


class ConfigWriteBreakpoint(gdb.Breakpoint):
    def __init__(self, address, image_base):
        super().__init__("*0x%x" % address, internal=False)
        self.image_base = image_base

    def stop(self):
        offset = int(gdb.parse_and_eval("$rsi"))
        value = int(gdb.parse_and_eval("$rdx")) & 0xffffffff
        if offset == 0x2070 or (GSP_ONLY and offset not in GSP_CONFIG_OFFSETS):
            return False
        caller = gdb.newest_frame().older().pc() - self.image_base
        gdb.write("config-write offset=%#x value=%#x caller=%#x\n" %
                  (offset, value, caller))
        return False


class ResolveConfigWrite(gdb.Breakpoint):
    def __init__(self, address, image_base):
        super().__init__("*0x%x" % address, internal=False, temporary=True)
        self.image_base = image_base

    def stop(self):
        self.enabled = False
        device = int(gdb.parse_and_eval("$rdi"))
        inferior = gdb.selected_inferior()
        pointer_size = gdb.lookup_type("void").pointer().sizeof
        vtable = int.from_bytes(inferior.read_memory(device, pointer_size),
                                byteorder="little")
        target = int.from_bytes(inferior.read_memory(vtable + 0x1F0,
                                                     pointer_size),
                                byteorder="little")
        gdb.write("config-write target=%#x library=%s\n" %
                  (target, gdb.solib_name(target)))
        ConfigWriteBreakpoint(target, self.image_base)
        return False


class ResolveRegisterTarget(gdb.Breakpoint):
    def __init__(self, address, register, label):
        super().__init__("*0x%x" % address, internal=False, temporary=True)
        self.register = register
        self.label = label

    def stop(self):
        self.enabled = False
        target = int(gdb.parse_and_eval("$" + self.register))
        gdb.write("%s target=%#x library=%s\n" %
                  (self.label, target, gdb.solib_name(target)))
        SyncBreakpoint(target, self.label,
                       [("this", "rdi"), ("buffer", "rsi"),
                        ("offset", "rdx"), ("bytes", "rcx"),
                        ("ops", "r8"), ("flags", "r9")])
        return False


def libpoplar_base():
    pid = gdb.selected_inferior().pid
    with open("/proc/%d/maps" % pid) as mappings:
        for line in mappings:
            if "libpoplar.so" in line and "r-xp" in line:
                start = int(line.split("-", 1)[0], 16)
                offset = int(line.split()[2], 16)
                return start - offset
    raise gdb.GdbError("libpoplar.so is not mapped")


base = libpoplar_base()
gdb.write("libpoplar base=%#x\n" % base)
ConfigWriteBreakpoint(base + 0x3C6C7A0, base)
if not GSP_ONLY:
    SyncBreakpoint(base + 0x3BD0040, "signal-execute", [("this", "rdi")])
    SyncBreakpoint(base + 0x3BFF360, "set-mark",
                   [("this", "rdi"), ("proxy", "rsi"), ("mark", "rdx")], base)
    SyncBreakpoint(base + 0x3BFE9A0, "wait-mark",
                   [("this", "rdi"), ("proxy", "rsi"), ("mark", "rdx"),
                    ("poll-us", "rcx"), ("timeout-us", "r8")])
    ResolveDeviceSetMark(base + 0x3BFF4B7)
    ResolveRegisterTarget(base + 0x3BC89B3, "r10", "mirror-buffer")
    ResolveRegisterTarget(base + 0x3BFD935, "rax", "device-wait-mark")
