From 19205f73a140d6d39b2d7855b7c1dffe0ef34266 Mon Sep 17 00:00:00 2001 From: Clement Rouault Date: Wed, 13 Jul 2016 18:23:26 +0200 Subject: [PATCH] Improve debugger handling of memBP for future use + fix MemBP prot init + add failing test of multitple BP with diff prot --- TODO | 1 + windows/debug/breakpoints.py | 2 +- windows/debug/debugger.py | 68 +++++++++++++++++------------------ windows/test/test_debugger.py | 49 +++++++++++++++++++++++++ 4 files changed, 85 insertions(+), 35 deletions(-) diff --git a/TODO b/TODO index da8fbfe..195471e 100644 --- a/TODO +++ b/TODO @@ -9,6 +9,7 @@ TODO: - Test !! (bp, BP_HX, bp on only on process, bp_hx on only one thread..) - test breakpoint with specific target - Add test for debugger with breakpoint that add another breakpoint on trigger + - Handle MemBP of multiple types of the same page.. - remotectypes - pretty sur I can get rid of PointerToStruct64/PointerToStruct32 diff --git a/windows/debug/breakpoints.py b/windows/debug/breakpoints.py index f7664b8..7e1e568 100644 --- a/windows/debug/breakpoints.py +++ b/windows/debug/breakpoints.py @@ -49,7 +49,7 @@ class MemoryBreakpoint(Breakpoint): def __init__(self, addr, size=None, prot=None): super(MemoryBreakpoint, self).__init__(addr) self.size = size if size is not None else self.DEFAULT_SIZE - self.protect = size if prot is not None else self.DEFAULT_PROTECT + self.protect = prot if prot is not None else self.DEFAULT_PROTECT def trigger(self, dbg, exception): diff --git a/windows/debug/debugger.py b/windows/debug/debugger.py index 79f944c..e1c3d61 100644 --- a/windows/debug/debugger.py +++ b/windows/debug/debugger.py @@ -1,5 +1,5 @@ import os.path -from collections import defaultdict +from collections import defaultdict, namedtuple from contextlib import contextmanager import windows @@ -30,6 +30,8 @@ class DEBUG_EVENT(DEBUG_EVENT): def code(self): return self.KNOWN_EVENT_CODE.get(self.dwDebugEventCode, self.dwDebugEventCode) +WatchedPage = namedtuple('WatchedPage', ["original_prot", "bps"]) + class Debugger(object): """A debugger based on standard Win32 API. Handle standard (int3) and Hardware-Exec Breakpoints""" @@ -241,12 +243,15 @@ class Debugger(object): old_prot = DWORD() vprot_begin = affected_pages[0] vprot_size = PAGE_SIZE * len(affected_pages) + print("[VP] {0:#x} {1:#x} {2}".format(vprot_begin, vprot_size, bp.protect)) target.virtual_protect(vprot_begin, vprot_size, bp.protect, old_prot) bp._old_prot = old_prot.value #self._virtual_protected_memory[vprot_begin] = (vprot_size, bp.protect, old_prot) cp_watch_page = self._watched_pages[self.current_process.pid] for page_addr in affected_pages: - cp_watch_page[page_addr].append(bp) + if page_addr not in cp_watch_page: + cp_watch_page[page_addr] = WatchedPage(old_prot, []) + cp_watch_page[page_addr].bps.append(bp) # TODO: watch for overlap with other MEM breakpoints return True @@ -263,8 +268,8 @@ class Debugger(object): cp_watch_page = self._watched_pages[self.current_process.pid] for page_addr in affected_pages: - cp_watch_page[page_addr].remove(bp) - if not cp_watch_page[page_addr]: + cp_watch_page[page_addr].bps.remove(bp) + if not cp_watch_page[page_addr].bps: del cp_watch_page[page_addr] else: raise NotImplementedError("Removing MemBP on page with multiple MemBP > 12) << 12 + cp_watch_page = self._watched_pages[self.current_process.pid] mem_bp = self.get_memory_breakpoint_at(fault_addr, self.current_process) if mem_bp is False: # No BP on this page return self.on_exception(exception) + original_prot = cp_watch_page[fault_page].original_prot if mem_bp is None: # Page as MEMBP but None handle this address # This hack is bad, find a BP on the page to restore original access.. - # TODO: stock original page protection elsewhere ? - bp = self._watched_pages[self.current_process.pid][fault_page][0] - self._pass_memory_breakpoint(bp, fault_page) + bp = cp_watch_page[fault_page].bps[-1] + self._pass_memory_breakpoint(bp, original_prot, fault_page) return DBG_CONTINUE continue_flag = mem_bp.trigger(self, exception) self._explicit_single_step[self.current_thread.tid] = self.current_thread.context.EEFlags.TF # If BP has not been removed in trigger, pas it - if mem_bp in self._watched_pages[self.current_process.pid][fault_page]: - self._pass_memory_breakpoint(mem_bp, fault_page) + if fault_page in cp_watch_page and mem_bp in cp_watch_page[fault_page].bps: + self._pass_memory_breakpoint(mem_bp, original_prot, fault_page) return continue_flag - #for bp, vprot_begin, vprot_end, original_prot in self._watched_memory: - # if vprot_begin <= fault_addr < vprot_end: - # # It's the page for this MEMBP that triggeed the BP - # if bp._addr <= fault_addr < bp._addr + bp.size: - # # In the real range of our memBP - # continue_flag = bp.trigger(self, exception) - # self._explicit_single_step[self.current_thread.tid] = self.current_thread.context.EEFlags.TF - # #if excp_addr in self.breakpoints[self.current_process.pid]: - # else: - # #self._explicit_single_step[self.current_thread.tid] = self.current_thread.context.EEFlags.TF - # continue_flag = DBG_CONTINUE - # - # if bp in [x[0] for x in self._watched_memory]: - # self._pass_memory_breakpoint(bp, vprot_begin, vprot_end, original_prot) - # return continue_flag - #else: - # self.on_exception(exception) - # TODO: self._explicit_single_step setup by single_step() ? check at the end ? finally ? def _handle_exception(self, debug_event): @@ -505,7 +505,7 @@ class Debugger(object): self._explicit_single_step[self.current_thread.tid] = False self._breakpoint_to_reput[self.current_thread.tid] = [] self.processes[self.current_process.pid] = self.current_process - self._watched_pages[self.current_process.pid] = defaultdict(list) + self._watched_pages[self.current_process.pid] = {} #defaultdict(list) self.breakpoints[self.current_process.pid] = {} self._module_by_process[self.current_process.pid] = {} self._update_debugger_state(debug_event) @@ -666,7 +666,7 @@ class Debugger(object): if fault_page not in self._watched_pages[process.pid]: return False - for bp in self._watched_pages[process.pid][fault_page]: + for bp in self._watched_pages[process.pid][fault_page].bps: if bp._addr <= addr < bp._addr + bp.size: return bp return None diff --git a/windows/test/test_debugger.py b/windows/test/test_debugger.py index bc759e7..6bf382b 100644 --- a/windows/test/test_debugger.py +++ b/windows/test/test_debugger.py @@ -614,6 +614,55 @@ class DebuggerTestCase(unittest.TestCase): d.loop() TEST_CASE.assertEqual(data, [data_addr, data_addr + 4]) + + def test_read_write_bp_same_page(self): + TEST_CASE = self + data = [] + + def generate_read_at(addr): + res = x86.MultipleInstr() + res += x86.Mov("EAX", x86.deref(addr)) + res += x86.Ret() + return res.get_code() + + def generate_write_at(addr): + res = x86.MultipleInstr() + res += x86.Mov(x86.deref(addr), "EAX") + res += x86.Ret() + return res.get_code() + + def do_check(): + calc.execute(generate_read_at(data_addr)).wait() + calc.execute(generate_write_at(data_addr + 4)).wait() + calc.execute(generate_read_at(data_addr + 0x500)).wait() + calc.execute(generate_write_at(data_addr + 0x504)).wait() + calc.exit() + + class MemBP(windows.debug.MemoryBreakpoint): + DEFAULT_PROTECT = PAGE_NOACCESS + def trigger(self, dbg, exc): + addr = exc.ExceptionRecord.ExceptionAddress + fault_addr = exc.ExceptionRecord.ExceptionInformation[1] + print("Got <{0:#x}> <{1}>".format(fault_addr, exc.ExceptionRecord.ExceptionInformation[0])) + data.append((self, fault_addr)) + + calc = pop_calc_32(dwCreationFlags=DEBUG_PROCESS) + d = windows.debug.Debugger(calc) + data_addr = calc.virtual_alloc(0x1000) + the_write_bp = MemBP(data_addr + 0x500, prot=PAGE_READONLY, size=0x500) + the_read_bp = MemBP(data_addr, prot=PAGE_NOACCESS, size=0x500) + d.add_bp(the_write_bp) + d.add_bp(the_read_bp) + threading.Thread(target=do_check).start() + d.loop() + + # generate_read_at (data_addr + 0x500)) (write_bp (PAGE_READONLY)) should not be triggered + expected_result = [(the_read_bp, data_addr), (the_read_bp, data_addr + 4), + (the_write_bp, data_addr + 0x504)] + + + TEST_CASE.assertEqual(data, expected_result) + if __name__ == '__main__': alltests = unittest.TestSuite() alltests.addTest(unittest.makeSuite(DebuggerTestCase))