diff --git a/droidasc/asc_client/asc_handler.py b/droidasc/asc_client/asc_handler.py index a44f1ef..63435ce 100644 --- a/droidasc/asc_client/asc_handler.py +++ b/droidasc/asc_client/asc_handler.py @@ -9,6 +9,7 @@ _DexManager = None _FindRefManager = None _decompile_dex_bytes = None +_decompile_dex_bytes_with_metadata = None _DEX = None @@ -36,7 +37,7 @@ def _install_pure_python_mutf8_shim(): def _lazy_import(): - global _DexManager, _FindRefManager, _decompile_dex_bytes, _DEX + global _DexManager, _FindRefManager, _decompile_dex_bytes, _decompile_dex_bytes_with_metadata, _DEX if _DexManager is not None: return @@ -44,11 +45,12 @@ def _lazy_import(): from droidasc.asc_core.core.dex.dex_manager import DexManager from droidasc.asc_core.findrefs.findrefs_manager import FindRefManager from droidasc.asc_core.utils.tinydex import DEX - from droidasc.asc_core.utils.decompiler import decompile_dex_bytes + from droidasc.asc_core.utils.decompiler import decompile_dex_bytes, decompile_dex_bytes_with_metadata _DexManager = DexManager _FindRefManager = FindRefManager _decompile_dex_bytes = decompile_dex_bytes + _decompile_dex_bytes_with_metadata = decompile_dex_bytes_with_metadata _DEX = DEX @@ -62,6 +64,12 @@ def getclass(self, dex_buf : bytes, dalvik_class : str) -> str: new_dex_bytes = manager.extract_and_rebuild(dalvik_class) return _decompile_dex_bytes(new_dex_bytes, dalvik_class) + def getclass_with_metadata(self, dex_buf : bytes, dalvik_class : str): + _lazy_import() + manager = _DexManager(memoryview(dex_buf), debug=self.debug) + new_dex_bytes = manager.extract_and_rebuild(dalvik_class) + return _decompile_dex_bytes_with_metadata(new_dex_bytes, dalvik_class) + def _format_method(self, dex, midx : int) -> str: method = dex.methods[midx] return f"{method.cls.fullname}->{method.name}" diff --git a/droidasc/asc_client/gui/app.py b/droidasc/asc_client/gui/app.py index 4714840..0979f04 100644 --- a/droidasc/asc_client/gui/app.py +++ b/droidasc/asc_client/gui/app.py @@ -13,10 +13,14 @@ identifier_occurrences_in_range, index_to_offset, is_identifier, + find_member_declaration, + linkable_member_spans, + member_reference_at_offset, + remap_ranges_after_replacements, rename_identifier_in_range, token_at_offset, ) -from droidasc.asc_client.gui.text_utils import decode_java_unicode_escapes +from droidasc.asc_client.gui.text_utils import decode_java_unicode_escapes_with_ranges from droidasc.asc_client.gui.theme import ThemeManager from droidasc.asc_client.gui.widgets import EditorTab, EditorTabBar from droidasc.asc_client.manifest_handler import get_manifest_xml @@ -242,6 +246,8 @@ def __init__(self, root, apk_path : str, max_workers : int = 8, debug : bool = F self._highlight_generation = 0 self._highlight_apply_batch = 400 self._highlight_apply_job = None + self._ctrl_member_links_visible = False + self._pending_member_navigation = None self._build_ui() self._bind_editor_shortcuts() @@ -301,6 +307,7 @@ def _configure_source_tags(self): self.source_text.tag_configure("symbol_current", background=t["symbol_current"]) self.source_text.tag_configure("find_match", background=t["find_match"], foreground=t["editor_fg"]) self.source_text.tag_configure("find_current", background=t["find_current"], foreground=t["editor_fg"]) + self.source_text.tag_configure("member_link", foreground=t["accent"], underline=True) def apply_theme(self): self.theme = self.theme_manager.theme @@ -450,6 +457,11 @@ def _bind_editor_shortcuts(self): self.root.bind("", self._next_tab) self.root.bind("", self._prev_tab) self.source_text.bind("", self._set_editor_cursor_from_click) + self.source_text.bind("", self._open_member_from_click) + self.root.bind("", self._show_member_links) + self.root.bind("", self._show_member_links) + self.root.bind("", self._hide_member_links) + self.root.bind("", self._hide_member_links) self.source_text.bind("", self._on_source_key_press) self.source_text.bind("", self._on_source_key_release) self.source_text.bind("<>", lambda _event: "break") @@ -494,8 +506,8 @@ def _handle_event(self, event): self.status_var.set(event[1]) return if kind == "source_done": - _kind, dalvik_class, dex_name, source = event - self._finish_open_tab(dalvik_class, dex_name, source) + _kind, dalvik_class, dex_name, source, member_references = event + self._finish_open_tab(dalvik_class, dex_name, source, member_references) return if kind == "text_tab_done": _kind, tab_id, title, source, dex_name = event @@ -688,11 +700,14 @@ def _set_source_text(self, text : str, reset_view : bool = True): def _render_tab_source(self, tab : EditorTab): source = tab.edited_source or tab.source if not tab.comments: + tab.rendered_member_references = list(tab.member_references) return source, [] out = [] comment_spans = [] + insertions = [] offset = 0 + source_offset = 0 lines = source.splitlines(keepends=True) for line_no, line in enumerate(lines, 1): newline = "" @@ -710,12 +725,14 @@ def _render_tab_source(self, tab : EditorTab): else: pad = _COMMENT_GAP comment_text = f"{pad}// {comment}" + insertions.append((source_offset + len(body), len(comment_text))) start = offset + len(body) + len(pad) end = start + len(comment_text) - len(pad) rendered = f"{body}{comment_text}{newline}" comment_spans.append((start, end)) out.append(rendered) offset += len(rendered) + source_offset += len(line) if not lines and tab.comments: comment = tab.comments.get(1, "") @@ -723,6 +740,16 @@ def _render_tab_source(self, tab : EditorTab): comment_spans.append((0, len(rendered))) return rendered, comment_spans + def shift_start(point): + return point + sum(length for position, length in insertions if position <= point) + + def shift_end(point): + return point + sum(length for position, length in insertions if position < point) + + tab.rendered_member_references = [ + (shift_start(item[0]), shift_end(item[1]), *item[2:]) + for item in tab.member_references + ] return "".join(out), comment_spans def _comment_line_from_index(self, index : str): @@ -752,9 +779,90 @@ def _set_editor_cursor_from_click(self, event): except tk.TclError: pass + def _show_member_links(self, _event = None): + if self._ctrl_member_links_visible: + return None + self._ctrl_member_links_visible = True + self.source_text.tag_remove("member_link", "1.0", tk.END) + tab = self._active_tab() + if tab is not None and tab.kind == "class" and not tab.loading and not tab.error: + self._tag_ranges( + "member_link", + linkable_member_spans( + tab.rendered_member_references, + self.store.class_to_dex if self.store is not None else (), + ), + ) + self.source_text.tag_raise("member_link") + return None + + def _hide_member_links(self, _event = None): + self._ctrl_member_links_visible = False + self.source_text.tag_remove("member_link", "1.0", tk.END) + return None + + def _member_at_index(self, index : str): + tab = self._active_tab() + if tab is None or tab.kind != "class" or tab.loading or tab.error: + return None + rendered = tab.edited_source or tab.source + references = tab.rendered_member_references or tab.member_references + if tab.comments: + rendered, _comment_spans = self._render_tab_source(tab) + references = tab.rendered_member_references + offset = index_to_offset(rendered, index) + return member_reference_at_offset(references, offset) + + def _open_member_from_click(self, event): + try: + index = self.source_text.index(f"@{event.x},{event.y}") + except tk.TclError: + return "break" + reference = self._member_at_index(index) + if reference is None: + return "break" + _start, _end, member_type, class_name, member_name, descriptor, _declaration = reference + if self.store is None or class_name not in self.store.class_to_dex: + self.status_var.set(f"Declaration class is not present in this APK: {dalvik_to_dot(class_name)}") + return "break" + self._pending_member_navigation = (class_name, member_type, member_name, descriptor) + self.open_class(class_name) + self._complete_pending_member_navigation(class_name) + return "break" + + def _complete_pending_member_navigation(self, dalvik_class : str): + target = self._pending_member_navigation + if target is None or target[0] != dalvik_class: + return + idx = self._find_tab_index(dalvik_class) + if idx is None: + return + tab = self.editor_tabs[idx] + if tab.loading or tab.error or not tab.source: + return + _class_name, member_type, member_name, descriptor = target + reference = find_member_declaration( + tab.rendered_member_references, member_type, member_name, descriptor + ) + self._pending_member_navigation = None + if reference is None: + self.status_var.set(f"Declaration not found: {dalvik_to_dot(dalvik_class)}->{member_name}") + return + start = reference[0] + index = f"1.0+{start}c" + self.source_text.mark_set("insert", index) + self.source_text.see(index) + tab.insert_index = self.source_text.index("insert") + tab.yview = self.source_text.yview() + self._highlight_active_line() + self._highlight_related_identifier(tab.insert_index) + self.status_var.set(f"Opened declaration {dalvik_to_dot(dalvik_class)}->{member_name}") + def _on_source_key_press(self, event): if event.state & 0x0004: return None + if event.keysym in ("x", "X"): + return self._find_current_member_references() if event.keysym == "n": return self._show_rename_dialog() if event.char == ";": @@ -769,6 +877,29 @@ def _on_source_key_release(self, event): self._highlight_related_identifier(self.source_text.index("insert")) return None + def _find_current_member_references(self, _event = None): + tab = self._active_tab() + if tab is None or tab.kind != "class" or tab.loading or tab.error or not tab.source: + return "break" + + try: + cursor_index = self.source_text.index("insert") + except tk.TclError: + return "break" + reference = self._member_at_index(cursor_index) + if reference is None: + self.status_var.set("Place the cursor on a method or field name to find its references") + return "break" + _start, _end, member_type, class_name, member_name, *_metadata = reference + + self.search_type_var.set(f"{member_type} refs") + self.search_value_var.set(member_name) + self.search_class_var.set(class_name) + self.fuzzy_class_var.set(False) + self._on_search_type_changed() + self._start_search(exact_member=True) + return "break" + def _highlight_related_identifier(self, index : str): self.source_text.tag_remove("symbol_match", "1.0", tk.END) self.source_text.tag_remove("symbol_current", "1.0", tk.END) @@ -776,6 +907,8 @@ def _highlight_related_identifier(self, index : str): if tab is None or tab.kind != "class" or tab.loading or tab.error: return source = tab.edited_source or tab.source + if tab.comments: + source, _comment_spans = self._render_tab_source(tab) if not source: return offset = index_to_offset(source, index) @@ -810,6 +943,7 @@ def _show_empty_editor(self): self.source_text.tag_remove("symbol_current", "1.0", tk.END) self.source_text.tag_remove("find_match", "1.0", tk.END) self.source_text.tag_remove("find_current", "1.0", tk.END) + self.source_text.tag_remove("member_link", "1.0", tk.END) self._set_source_text("Open a class from the package tree or search results.", reset_view=True) self.editor_find_status_var.set("") @@ -822,6 +956,7 @@ def _show_tab(self, tab : EditorTab): self.source_text.tag_remove("symbol_current", "1.0", tk.END) self.source_text.tag_remove("find_match", "1.0", tk.END) self.source_text.tag_remove("find_current", "1.0", tk.END) + self.source_text.tag_remove("member_link", "1.0", tk.END) if tab.loading: self._set_source_text(f"Decompiling {dalvik_to_dot(tab.dalvik_class)}...", reset_view=True) return @@ -845,6 +980,9 @@ def _show_tab(self, tab : EditorTab): self._highlight_quoted_strings(rendered) self._refresh_editor_find_marks(reset_cursor=False) self._highlight_active_line() + if self._ctrl_member_links_visible: + self._ctrl_member_links_visible = False + self._show_member_links() def activate_tab(self, index : int): if index < 0 or index >= len(self.editor_tabs): @@ -852,6 +990,7 @@ def activate_tab(self, index : int): if index == self.active_tab_idx: self.editor_tab_bar.ensure_visible(index) self.editor_tab_bar.redraw() + self._complete_pending_member_navigation(self.editor_tabs[index].dalvik_class) return self._remember_active_view() self.active_tab_idx = index @@ -870,13 +1009,20 @@ def activate_tab(self, index : int): self.status_var.set(f"Opened {tab.title}") if tab.kind == "class": self._select_class_in_tree(tab.dalvik_class) + self._complete_pending_member_navigation(tab.dalvik_class) def close_tab(self, index : int): if index < 0 or index >= len(self.editor_tabs): return self._remember_active_view() + closing_tab = self.editor_tabs[index] was_active = index == self.active_tab_idx self.editor_tabs.pop(index) + if ( + self._pending_member_navigation is not None + and self._pending_member_navigation[0] == closing_tab.dalvik_class + ): + self._pending_member_navigation = None if not self.editor_tabs: self.active_tab_idx = -1 self.editor_tab_bar.tab_scroll_index = 0 @@ -897,6 +1043,7 @@ def clear_tabs(self): return self.editor_tabs = [] self.active_tab_idx = -1 + self._pending_member_navigation = None self.editor_tab_bar.tab_scroll_index = 0 self.editor_tab_bar.redraw() self._show_empty_editor() @@ -918,13 +1065,15 @@ def _next_tab(self, _event = None): def _prev_tab(self, _event = None): return self._cycle_tab(-1) - def _finish_open_tab(self, dalvik_class : str, dex_name : str, source : str): + def _finish_open_tab(self, dalvik_class : str, dex_name : str, source : str, member_references): idx = self._find_tab_index(dalvik_class) if idx is None: return tab = self.editor_tabs[idx] tab.dex_name = dex_name - tab.source = decode_java_unicode_escapes(source) + tab.source, tab.member_references = decode_java_unicode_escapes_with_ranges( + source, member_references + ) tab.edited_source = "" tab.loading = False tab.error = "" @@ -933,6 +1082,8 @@ def _finish_open_tab(self, dalvik_class : str, dex_name : str, source : str): self.editor_tab_bar.ensure_visible(idx if idx == self.active_tab_idx else self.active_tab_idx) self.editor_tab_bar.redraw() self.status_var.set(f"Opened {dalvik_to_dot(dalvik_class)} from {dex_name}") + if idx == self.active_tab_idx: + self._complete_pending_member_navigation(dalvik_class) def _fail_open_tab(self, dalvik_class : str, message : str): idx = self._find_tab_index(dalvik_class) @@ -942,6 +1093,11 @@ def _fail_open_tab(self, dalvik_class : str, message : str): tab = self.editor_tabs[idx] tab.loading = False tab.error = message + if ( + self._pending_member_navigation is not None + and self._pending_member_navigation[0] == dalvik_class + ): + self._pending_member_navigation = None if idx == self.active_tab_idx: self._show_tab(tab) self.editor_tab_bar.redraw() @@ -1417,6 +1573,9 @@ def _rename_identifier(self, tab : EditorTab, method_range, old_name : str, new_ if old_name == new_name: return source = tab.edited_source or tab.source + renamed_spans = identifier_occurrences_in_range( + source, method_range[0], method_range[1], old_name + ) new_source, count = rename_identifier_in_range( source, method_range[0], @@ -1428,6 +1587,9 @@ def _rename_identifier(self, tab : EditorTab, method_range, old_name : str, new_ self.status_var.set(f"No occurrences of {old_name} in current method") return tab.edited_source = new_source + tab.member_references = remap_ranges_after_replacements( + tab.member_references, renamed_spans, len(new_name) + ) try: tab.insert_index = self.source_text.index("insert") tab.yview = self.source_text.yview() @@ -1465,8 +1627,10 @@ def open_class(self, dalvik_class : str): def worker(): try: _debug_log(self.debug, "app", f"open class worker start class={dalvik_class}") - dex_name, source = self.store.get_source(dalvik_class) - self.events.put(("source_done", dalvik_class, dex_name, source)) + dex_name, source, member_references = self.store.get_source_with_metadata(dalvik_class) + self.events.put(( + "source_done", dalvik_class, dex_name, source, member_references + )) _debug_log(self.debug, "app", f"open class worker done class={dalvik_class} dex={dex_name}") except Exception as e: if self.debug: @@ -1475,7 +1639,7 @@ def worker(): threading.Thread(target=worker, daemon=True).start() - def _start_search(self): + def _start_search(self, exact_member : bool = False): if self.store is None: return if self.search_inflight: @@ -1531,6 +1695,7 @@ def worker(): value=value, class_name=class_name or None, fuzzy_class=fuzzy_class, + exact_member=exact_member, max_workers=self.max_workers, progress_callback=lambda done, total, hit_count: self.events.put( ("search_progress", done, total, hit_count) diff --git a/droidasc/asc_client/gui/runtime.py b/droidasc/asc_client/gui/runtime.py index e916038..3047061 100644 --- a/droidasc/asc_client/gui/runtime.py +++ b/droidasc/asc_client/gui/runtime.py @@ -81,7 +81,10 @@ def _normalize_class_query(name : str, fuzzy : bool): return format_class_name(name) -def build_find_query(find_type : str, value : str, class_name = None, fuzzy_class : bool = False): +def build_find_query( + find_type : str, value : str, class_name = None, fuzzy_class : bool = False, + exact_member : bool = False, +): find_type = _REF_SEARCH_TYPES.get(find_type, find_type) if find_type == "string": return "string", {"string": value} @@ -92,9 +95,10 @@ def build_find_query(find_type : str, value : str, class_name = None, fuzzy_clas if class_name is None and not value: raise ValueError(f"{find_type} query needs at least one of class or {find_type} name") + member = [value, True] if exact_member and value else value or None if class_name is None: - return find_type, {find_type: {"class": None, find_type: value or None}} - return find_type, {find_type: {"class": [class_name, not fuzzy_class], find_type: value or None}} + return find_type, {find_type: {"class": None, find_type: member}} + return find_type, {find_type: {"class": [class_name, not fuzzy_class], find_type: member}} def parse_result_line(line : str): @@ -238,6 +242,10 @@ def iter_filtered_classes(self, keyword : str, limit : int): return ret def get_source(self, dalvik_class : str): + dex_name, source, _references = self.get_source_with_metadata(dalvik_class) + return dex_name, source + + def get_source_with_metadata(self, dalvik_class : str): with self._source_lock: old = self.source_cache.get(dalvik_class) if old is not None: @@ -259,10 +267,10 @@ def get_source(self, dalvik_class : str): with _GUI_MP_LOCK: snapshot = _snapshot_sys_modules() try: - source = AscHandler(self.debug).getclass(dex_buf, dalvik_class) + source, references = AscHandler(self.debug).getclass_with_metadata(dex_buf, dalvik_class) finally: _restore_sys_modules(snapshot) - ret = (dex_name, source) + ret = (dex_name, source, references) with self._source_lock: self.source_cache[dalvik_class] = ret _debug_log( @@ -505,11 +513,12 @@ def search( value : str, class_name = None, fuzzy_class : bool = False, + exact_member : bool = False, max_workers = None, progress_callback = None, result_callback = None, ): - find_type, find = build_find_query(find_type, value, class_name, fuzzy_class) + find_type, find = build_find_query(find_type, value, class_name, fuzzy_class, exact_member) workers = self.get_effective_search_workers(max_workers) _debug_log( self.debug, diff --git a/droidasc/asc_client/gui/source_edit.py b/droidasc/asc_client/gui/source_edit.py index 7c2ecba..e916396 100644 --- a/droidasc/asc_client/gui/source_edit.py +++ b/droidasc/asc_client/gui/source_edit.py @@ -245,6 +245,49 @@ def rename_identifier_in_range(text : str, start : int, end : int, old : str, ne return text[:start] + "".join(out) + text[end:], len(spans) +def remap_ranges_after_replacements(ranges, spans, replacement_length : int): + def remap(point): + shift = 0 + for start, end in spans: + if point < end: + break + shift += replacement_length - (end - start) + return point + shift + + return [ + (remap(item[0]), remap(item[1]), *item[2:]) + for item in ranges + ] + + +def member_reference_at_offset(references, offset : int): + return next( + (item for item in references if item[0] <= offset < item[1]), + None, + ) + + +def find_member_declaration(references, member_type : str, member_name : str, descriptor : str): + declarations = [ + item for item in references + if item[2] == member_type and item[4] == member_name and item[6] + ] + exact = next((item for item in declarations if item[5] == descriptor), None) + if exact is not None: + return exact + if not descriptor and declarations: + return declarations[0] + return None + + +def linkable_member_spans(references, available_classes): + return [ + (item[0], item[1]) + for item in references + if item[3] in available_classes + ] + + def identifier_occurrences_in_range(text : str, start : int, end : int, name : str): spans = [] state = "code" diff --git a/droidasc/asc_client/gui/text_utils.py b/droidasc/asc_client/gui/text_utils.py index 5bef690..0368dba 100644 --- a/droidasc/asc_client/gui/text_utils.py +++ b/droidasc/asc_client/gui/text_utils.py @@ -79,3 +79,24 @@ def decode_java_unicode_escapes(text : str): if pending_high is not None: out.append(pending_high_raw) return "".join(out) + + +def decode_java_unicode_escapes_with_ranges(text : str, ranges): + if not ranges or "\\u" not in text: + return decode_java_unicode_escapes(text), list(ranges) + + points = sorted({point for item in ranges for point in item[:2]}) + mapped = {} + source_pos = 0 + decoded_pos = 0 + for point in points: + decoded_pos += len(decode_java_unicode_escapes(text[source_pos:point])) + mapped[point] = decoded_pos + source_pos = point + + decoded = decode_java_unicode_escapes(text) + adjusted = [ + (mapped[item[0]], mapped[item[1]], *item[2:]) + for item in ranges + ] + return decoded, adjusted diff --git a/droidasc/asc_client/gui/widgets.py b/droidasc/asc_client/gui/widgets.py index 944421e..5fc35a1 100644 --- a/droidasc/asc_client/gui/widgets.py +++ b/droidasc/asc_client/gui/widgets.py @@ -20,6 +20,8 @@ class EditorTab: xview : tuple = (0.0, 1.0) insert_index : str = "1.0" comments : dict = field(default_factory=dict) + member_references : list = field(default_factory=list) + rendered_member_references : list = field(default_factory=list) class EditorTabBar(ttk.Frame): diff --git a/droidasc/asc_core/findrefs/locator/field_locator.py b/droidasc/asc_core/findrefs/locator/field_locator.py index 2bb4a64..e68e588 100644 --- a/droidasc/asc_core/findrefs/locator/field_locator.py +++ b/droidasc/asc_core/findrefs/locator/field_locator.py @@ -70,11 +70,12 @@ def _collect_clz_fids(self, type_idxs) -> set: ret.update(fids) return ret - def _match_clz_fids(self, clz_fids : set, field : str) -> set: + def _match_clz_fids(self, clz_fids : set, field : str, precise : bool = False) -> set: ret = set() fields = self.dex.fields for fid in clz_fids: - if fields[fid].name.find(field) != -1: + matches = fields[fid].name == field if precise else field in fields[fid].name + if matches: ret.add(fid) return ret @@ -91,6 +92,9 @@ def locate(self, find : dict) -> set: clz = find.get("class") field = find.get("field") + field_precise = False + if isinstance(field, (list, tuple)): + field, field_precise = field clz_precise = True if clz is not None: clz, clz_precise = clz @@ -103,6 +107,8 @@ def locate(self, find : dict) -> set: return set() if clz is None: + if field_precise: + return {fid for fid, item in enumerate(self.dex.fields) if item.name == field} name_idxs = self.str_locator.locate(field) field_maps = self.field_maps ret = set() @@ -133,10 +139,12 @@ def locate(self, find : dict) -> set: self._debug_log("locate", t_start, len(clz_fids)) return clz_fids if clz_precise: - ret = self._match_clz_fids(clz_fids, field) + ret = self._match_clz_fids(clz_fids, field, field_precise) self._debug_log("locate", t_start, len(ret)) return ret + if field_precise: + return self._match_clz_fids(clz_fids, field, True) name_idxs = self.str_locator.locate(field) ret = set() for name_idx in name_idxs: diff --git a/droidasc/asc_core/findrefs/locator/method_locator.py b/droidasc/asc_core/findrefs/locator/method_locator.py index cbadc1a..ed85f58 100644 --- a/droidasc/asc_core/findrefs/locator/method_locator.py +++ b/droidasc/asc_core/findrefs/locator/method_locator.py @@ -69,11 +69,12 @@ def _collect_clz_mids(self, type_idxs) -> set: ret.update(mids) return ret - def _match_clz_mids(self, clz_mids : set, method : str) -> set: + def _match_clz_mids(self, clz_mids : set, method : str, precise : bool = False) -> set: ret = set() methods = self.dex.methods for mid in clz_mids: - if methods[mid].name.find(method) != -1: + matches = methods[mid].name == method if precise else method in methods[mid].name + if matches: ret.add(mid) return ret @@ -90,6 +91,9 @@ def locate(self, find : dict) -> set: clz = find.get("class") method = find.get("method") + method_precise = False + if isinstance(method, (list, tuple)): + method, method_precise = method clz_precise = True if clz is not None: clz, clz_precise = clz @@ -102,6 +106,8 @@ def locate(self, find : dict) -> set: return set() if clz is None: + if method_precise: + return {mid for mid, item in enumerate(self.dex.methods) if item.name == method} name_idxs = self.str_locator.locate(method) method_maps = self.method_maps ret = set() @@ -132,10 +138,12 @@ def locate(self, find : dict) -> set: self._debug_log("locate", t_start, len(clz_mids)) return clz_mids if clz_precise: - ret = self._match_clz_mids(clz_mids, method) + ret = self._match_clz_mids(clz_mids, method, method_precise) self._debug_log("locate", t_start, len(ret)) return ret + if method_precise: + return self._match_clz_mids(clz_mids, method, True) name_idxs = self.str_locator.locate(method) ret = set() for name_idx in name_idxs: diff --git a/droidasc/asc_core/utils/decompiler.py b/droidasc/asc_core/utils/decompiler.py index 02623c9..bc9134a 100644 --- a/droidasc/asc_core/utils/decompiler.py +++ b/droidasc/asc_core/utils/decompiler.py @@ -95,6 +95,7 @@ def monkey_header_init(self, offset, buff, cm): androguard_dex.HeaderItem.__init__ = monkey_header_init from androguard.decompiler import decompile +from androguard.decompiler import instruction as androguard_instruction from androguard.decompiler import util as androguard_util from androguard.core.analysis.analysis import MethodAnalysis import androguard.core.androconf as androconf @@ -225,6 +226,66 @@ def patched_dvmethod_get_source(self) -> str: return "".join("\n %s" % annotation for annotation in annotations) + source decompile.DvMethod.get_source = patched_dvmethod_get_source +_orig_dvmethod_get_source_ext = decompile.DvMethod.get_source_ext +def patched_dvmethod_get_source_ext(self) -> list[tuple]: + source = _orig_dvmethod_get_source_ext(self) + annotations = getattr(self, "_asc_annotations", None) + if not annotations: + return source + prefix = "".join("\n %s" % annotation for annotation in annotations) + return [("ASC_ANNOTATIONS", prefix)] + source + +decompile.DvMethod.get_source_ext = patched_dvmethod_get_source_ext + +# Preserve the declaring field on source tokens. Androguard's default Writer +# keeps the field name for instance accesses and only a combined text token for +# static accesses, but drops the field IR object containing clsdesc. +def patched_instance_expression_visit(self, visitor): + return visitor.visit_get_instance( + self.var_map[self.arg], self.name, data=self + ) + + +def patched_instance_instruction_visit(self, visitor): + v_m = self.var_map + return visitor.visit_put_instance( + v_m[self.lhs], self.name, v_m[self.rhs], data=self + ) + + +def patched_writer_get_static(self, cls, name, data=None): + value = f"{cls}.{name}" + self.write(value) + self.write_ext(("GET_STATIC", value, data)) + + +def patched_static_expression_visit(self, visitor): + return visitor.visit_get_static(self.cls, self.name, data=self) + + +def patched_writer_put_static(self, cls, name, rhs, data=None): + self.write_ind() + value = f"{cls}.{name}" + self.write(value) + self.write_ext(("PUT_STATIC", value, data)) + self.write(" = ") + self.write_ext(("FIELD_ASSIGN", " = ")) + rhs.visit(self) + self.end_ins() + + +def patched_static_instruction_visit(self, visitor): + return visitor.visit_put_static( + self.cls, self.name, self.var_map[self.rhs], data=self + ) + + +androguard_instruction.InstanceExpression.visit = patched_instance_expression_visit +androguard_instruction.InstanceInstruction.visit = patched_instance_instruction_visit +androguard_instruction.StaticExpression.visit = patched_static_expression_visit +androguard_instruction.StaticInstruction.visit = patched_static_instruction_visit +decompile.Writer.visit_get_static = patched_writer_get_static +decompile.Writer.visit_put_static = patched_writer_put_static # --- End of String/Type Optimizations --- # We also completely disable ALL androguard loggers via python's standard logging module @@ -265,7 +326,7 @@ def get_method(self, method): self.methods[method] = ma return self.methods[method] -def decompile_dex_bytes(dex_bytes: bytearray, dalvik_class_fmt: str): +def _decompile_class(dex_bytes: bytearray, dalvik_class_fmt: str): """ Take DEX bytes and a target class format, decompile it using Androguard DAD and return the source code. @@ -277,9 +338,104 @@ def decompile_dex_bytes(dex_bytes: bytearray, dalvik_class_fmt: str): target_class = d.get_class(dalvik_class_fmt) if not target_class: - return f"Error: Class {dalvik_class_fmt} not found in the reconstructed DEX." + return None c = decompile.DvClass(target_class, dx) c.process() + return c + + +def _dalvik_class_name(name : str) -> str: + name = (name or "").replace(".", "/") + if not name.startswith("L"): + name = "L" + name + if not name.endswith(";"): + name += ";" + return name + + +# reuse androguard daddecompiler's token type, we can extract source' semantics, enhance gui xref ability +# 20260921 we may extract daddecompiler from androguard in future, for better development +def _source_with_member_references(source_ext): + parts = [] + references = [] + offset = 0 + + def append_tokens(tokens): + nonlocal offset + for token in tokens: + if len(token) < 2: + continue + kind, value = token[0], token[1] + if isinstance(value, list): + append_tokens(value) + continue + value = str(value) + start = offset + parts.append(value) + offset += len(value) + + if kind == "NAME_METHOD_PROTOTYPE" and len(token) >= 3: + method = token[2] + references.append(( + start, offset, "method", + _dalvik_class_name(method.cls_name), + method.name, + method.triple[2], + True, + )) + elif kind == "NAME_METHOD_INVOKE" and len(token) >= 7: + invoke = token[6] + references.append(( + start, offset, "method", + _dalvik_class_name(invoke.triple[0]), + invoke.name, + invoke.triple[2], + False, + )) + elif kind == "NAME_FIELD" and len(token) >= 4: + field = token[3] + references.append(( + start, offset, "field", + field.get_class_name(), + field.get_name(), + field.get_descriptor(), + True, + )) + elif kind in ("NAME_CLASS_INSTANCE", "NAME_CLASS_ASSIGNMENT") and len(token) >= 3: + field = token[2] + references.append(( + start, offset, "field", + _dalvik_class_name(field.clsdesc), + field.name, + getattr(field, "ftype", getattr(field, "atype", "")), + False, + )) + elif kind in ("GET_STATIC", "PUT_STATIC") and len(token) >= 3: + field = token[2] + field_start = offset - len(field.name) + references.append(( + field_start, offset, "field", + _dalvik_class_name(field.clsdesc), + field.name, + field.ftype, + False, + )) + + append_tokens(source_ext) + return "".join(parts), references + + +def decompile_dex_bytes(dex_bytes: bytearray, dalvik_class_fmt: str): + c = _decompile_class(dex_bytes, dalvik_class_fmt) + if c is None: + return f"Error: Class {dalvik_class_fmt} not found in the reconstructed DEX." # Remove the Decompile only time debug output return c.get_source() + + +def decompile_dex_bytes_with_metadata(dex_bytes: bytearray, dalvik_class_fmt: str): + c = _decompile_class(dex_bytes, dalvik_class_fmt) + if c is None: + return f"Error: Class {dalvik_class_fmt} not found in the reconstructed DEX.", [] + return _source_with_member_references(c.get_source_ext()) diff --git a/tests/dex_fixture.py b/tests/dex_fixture.py index 30417a7..cdfbf2e 100644 --- a/tests/dex_fixture.py +++ b/tests/dex_fixture.py @@ -56,6 +56,167 @@ def section(kind, size, data, align=4): return bytes(buf) +def make_invoke_dex(): + """Build a class whose run method invokes example.Target.callToMethod.""" + strings = [ + b'Lexample/Test;', b'Ljava/lang/Object;', b'V', b'run', + b'Lexample/Target;', b'callToMethod', + ] + buf = bytearray(112) + sections = [(0, 1, 0)] + + def section(kind, size, data, align=4): + buf.extend(b'\0' * (-len(buf) % align)) + off = len(buf) + sections.append((kind, size, off)) + buf.extend(data) + return off + + string_ids = section(1, len(strings), bytes(4 * len(strings))) + type_ids = section(2, 4, struct.pack('