mirror of
https://github.com/cyber-defence-campus/mole
synced 2026-06-20 13:19:21 +00:00
182 increased memory usage during path analysis (#185)
* Limit cache entries * Use breadth-first-serach and allow to limit visited memory versions * Add setting max_memory_slice_depth * Limit cache entries * Use breadth-first-serach and allow to limit visited memory versions * Add setting max_memory_slice_depth
This commit is contained in:
@@ -50,6 +50,12 @@ def main() -> None:
|
||||
default=None,
|
||||
help="maximum slice depth to stop the search",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_memory_slice_depth",
|
||||
type=int,
|
||||
default=None,
|
||||
help="maximum memory slice depth to stop the search",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--export_paths_to_json_file", help="export identified paths in JSON format"
|
||||
)
|
||||
@@ -71,6 +77,7 @@ def main() -> None:
|
||||
max_workers=args["max_workers"],
|
||||
max_call_level=args["max_call_level"],
|
||||
max_slice_depth=args["max_slice_depth"],
|
||||
max_memory_slice_depth=args["max_memory_slice_depth"],
|
||||
)
|
||||
slicer.start()
|
||||
paths = slicer.paths()
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from __future__ import annotations
|
||||
from collections import deque
|
||||
from functools import lru_cache
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
import binaryninja as bn
|
||||
@@ -61,7 +62,7 @@ class FunctionHelper:
|
||||
return parm_insts
|
||||
|
||||
@staticmethod
|
||||
@lru_cache(maxsize=None)
|
||||
@lru_cache(maxsize=32)
|
||||
def get_var_addr_assignments(
|
||||
func: bn.MediumLevelILFunction,
|
||||
) -> Dict[bn.Variable, List[bn.MediumLevelILSetVarSsa]]:
|
||||
@@ -97,39 +98,44 @@ class FunctionHelper:
|
||||
return var_addr_assignments
|
||||
|
||||
@staticmethod
|
||||
@lru_cache(maxsize=None)
|
||||
@lru_cache(maxsize=32)
|
||||
def get_ssa_memory_definitions(
|
||||
func: bn.MediumLevelILFunction,
|
||||
ssa_memory_version: int,
|
||||
ssa_memory_versions: frozenset[int] = frozenset(),
|
||||
memory_version: int,
|
||||
max_memory_slice_depth: int = -1,
|
||||
) -> List[bn.MediumLevelILInstruction]:
|
||||
"""
|
||||
This method returns a list of all instructions within `func` that define the memory with
|
||||
version `ssa_memory_version`. A memory defining instruction is an instruction that creates a
|
||||
new memory version.
|
||||
version `memory_version` (using breadth-first search). A memory defining instruction is an
|
||||
instruction that creates a new memory version.
|
||||
"""
|
||||
# Determine current memory defining instruction
|
||||
if func is None or ssa_memory_version in ssa_memory_versions:
|
||||
if func is None:
|
||||
return []
|
||||
ssa_memory_versions = ssa_memory_versions.union({ssa_memory_version})
|
||||
mem_def_inst = func.get_ssa_memory_definition(ssa_memory_version)
|
||||
if mem_def_inst is None:
|
||||
return []
|
||||
mem_def_insts: List[bn.MediumLevelILInstruction] = [mem_def_inst]
|
||||
# Determine source memory versions
|
||||
src_memory_versions: List[int] = []
|
||||
match mem_def_inst:
|
||||
case bn.MediumLevelILMemPhi(src_memory=src_memory):
|
||||
src_memory_versions.extend(src_memory)
|
||||
case _:
|
||||
src_memory_versions.append(mem_def_inst.ssa_memory_version)
|
||||
# Recursively determine memory defining instructions
|
||||
for src_memory_version in src_memory_versions:
|
||||
mem_def_insts.extend(
|
||||
FunctionHelper.get_ssa_memory_definitions(
|
||||
func, src_memory_version, ssa_memory_versions
|
||||
)
|
||||
)
|
||||
mem_def_insts: List[bn.MediumLevelILInstruction] = []
|
||||
visited_memory_versions = set()
|
||||
queue = deque([memory_version])
|
||||
while queue:
|
||||
# Break if maximum number of memory versions visited
|
||||
if (
|
||||
max_memory_slice_depth >= 0
|
||||
and len(visited_memory_versions) >= max_memory_slice_depth
|
||||
):
|
||||
break
|
||||
# Get current memory version
|
||||
current_memory_version = queue.popleft()
|
||||
if current_memory_version not in visited_memory_versions:
|
||||
# Visit current memory version
|
||||
visited_memory_versions.add(current_memory_version)
|
||||
mem_def_inst = func.get_ssa_memory_definition(current_memory_version)
|
||||
if mem_def_inst is None:
|
||||
continue
|
||||
mem_def_insts.append(mem_def_inst)
|
||||
# Add new memory versions to queue
|
||||
match mem_def_inst:
|
||||
case bn.MediumLevelILMemPhi(src_memory=src_memory):
|
||||
queue.extend(src_memory)
|
||||
case _:
|
||||
queue.append(mem_def_inst.ssa_memory_version)
|
||||
return mem_def_insts
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -10,7 +10,7 @@ class InstructionHelper:
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
@lru_cache(maxsize=None)
|
||||
@lru_cache(maxsize=64)
|
||||
def format_inst(inst: bn.MediumLevelILInstruction) -> str:
|
||||
"""
|
||||
This method replaces function addresses with their names.
|
||||
|
||||
@@ -17,6 +17,12 @@ settings:
|
||||
value: -1
|
||||
min_value: -1
|
||||
max_value: 9999
|
||||
max_memory_slice_depth:
|
||||
name: max_memory_slice_depth
|
||||
help: maximum memory slice depth to stop the search
|
||||
value: -1
|
||||
min_value: -1
|
||||
max_value: 9999
|
||||
src_highlight_color:
|
||||
name: src_highlight_color
|
||||
help: color used to highlight instructions originating from slicing a source function
|
||||
|
||||
+7
-2
@@ -326,7 +326,7 @@ class SourceFunction(Function):
|
||||
continue
|
||||
# Create backward slicer
|
||||
src_slicer = MediumLevelILBackwardSlicer(
|
||||
bv, custom_tag, 0, cancelled
|
||||
bv, custom_tag, 0, 0, cancelled
|
||||
)
|
||||
# Add edge between call and parameter instructions
|
||||
src_slicer.inst_graph.add_node(
|
||||
@@ -389,6 +389,7 @@ class SinkFunction(Function):
|
||||
manual_fun_all_code_xrefs: bool,
|
||||
max_call_level: int,
|
||||
max_slice_depth: int,
|
||||
max_memory_slice_depth: int,
|
||||
found_path: Callable[[Path], None],
|
||||
cancelled: Callable[[], bool],
|
||||
) -> List[Path]:
|
||||
@@ -513,7 +514,11 @@ class SinkFunction(Function):
|
||||
if par_slice_fun(snk_par_idx):
|
||||
# Create backward slicer
|
||||
snk_slicer = MediumLevelILBackwardSlicer(
|
||||
bv, custom_tag, max_call_level, cancelled
|
||||
bv,
|
||||
custom_tag,
|
||||
max_call_level,
|
||||
max_memory_slice_depth,
|
||||
cancelled,
|
||||
)
|
||||
snk_inst_graph = snk_slicer.inst_graph
|
||||
snk_call_graph = snk_slicer.call_graph
|
||||
|
||||
+8
-2
@@ -169,6 +169,7 @@ class MediumLevelILBackwardSlicer:
|
||||
bv: bn.BinaryView,
|
||||
custom_tag: str = "",
|
||||
max_call_level: int = -1,
|
||||
max_memory_slice_depth: int = -1,
|
||||
cancelled: Callable[[], bool] = None,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -182,6 +183,7 @@ class MediumLevelILBackwardSlicer:
|
||||
elif "snk" in self._tag.lower():
|
||||
self._origin = "snk"
|
||||
self._max_call_level: int = max_call_level
|
||||
self._max_memory_slice_depth: int = max_memory_slice_depth
|
||||
self._cancelled = cancelled
|
||||
self._inst_visited: Set[bn.MediumLevelILInstruction] = set()
|
||||
self.inst_graph: MediumLevelILInstructionGraph = MediumLevelILInstructionGraph()
|
||||
@@ -338,7 +340,9 @@ class MediumLevelILBackwardSlicer:
|
||||
case bn.MediumLevelILConstPtr():
|
||||
# Iterate all memory defining instructions
|
||||
mem_def_insts = FunctionHelper.get_ssa_memory_definitions(
|
||||
inst.function, inst.ssa_memory_version
|
||||
inst.function,
|
||||
inst.ssa_memory_version,
|
||||
self._max_memory_slice_depth,
|
||||
)
|
||||
for mem_def_inst in mem_def_insts:
|
||||
mem_def_inst_info = InstructionHelper.get_inst_info(
|
||||
@@ -426,7 +430,9 @@ class MediumLevelILBackwardSlicer:
|
||||
dest_var_use_sites[dest_var_use_site] = var_addr_ass_inst
|
||||
# Get all instructions in the current function defining the current memory version
|
||||
mem_def_insts = FunctionHelper.get_ssa_memory_definitions(
|
||||
inst.function, inst.ssa_memory_version
|
||||
inst.function,
|
||||
inst.ssa_memory_version,
|
||||
self._max_memory_slice_depth,
|
||||
)
|
||||
for mem_def_inst in mem_def_insts:
|
||||
mem_def_inst_info = InstructionHelper.get_inst_info(
|
||||
|
||||
@@ -196,6 +196,7 @@ class ConfigService:
|
||||
"max_workers",
|
||||
"max_call_level",
|
||||
"max_slice_depth",
|
||||
"max_memory_slice_depth",
|
||||
"max_turns",
|
||||
"max_completion_tokens",
|
||||
]:
|
||||
|
||||
@@ -24,6 +24,7 @@ class PathService(BackgroundTask):
|
||||
max_workers: Optional[int] = None,
|
||||
max_call_level: Optional[int] = None,
|
||||
max_slice_depth: Optional[int] = None,
|
||||
max_memory_slice_depth: Optional[int] = None,
|
||||
enable_all_funs: bool = False,
|
||||
manual_fun: Optional[SourceFunction | SinkFunction] = None,
|
||||
manual_fun_inst: Optional[
|
||||
@@ -46,6 +47,7 @@ class PathService(BackgroundTask):
|
||||
self._max_workers = max_workers
|
||||
self._max_call_level = max_call_level
|
||||
self._max_slice_depth = max_slice_depth
|
||||
self._max_memory_slice_depth = max_memory_slice_depth
|
||||
self._enable_all_funs = enable_all_funs
|
||||
self._manual_fun = manual_fun
|
||||
self._manual_fun_inst = manual_fun_inst
|
||||
@@ -92,6 +94,12 @@ class PathService(BackgroundTask):
|
||||
if setting:
|
||||
max_slice_depth = setting.value
|
||||
log.debug(tag, f"- max_slice_depth: '{max_slice_depth}'")
|
||||
max_memory_slice_depth = self._max_memory_slice_depth
|
||||
if max_memory_slice_depth is None:
|
||||
setting = self._config_model.get_setting("max_memory_slice_depth")
|
||||
if setting:
|
||||
max_memory_slice_depth = setting.value
|
||||
log.debug(tag, f"- max_memory_slice_depth: '{max_memory_slice_depth}'")
|
||||
# Source functions
|
||||
src_funs: List[SourceFunction] = self._config_model.get_functions(
|
||||
fun_type="Sources",
|
||||
@@ -177,6 +185,7 @@ class PathService(BackgroundTask):
|
||||
self._manual_fun_all_code_xrefs,
|
||||
max_call_level,
|
||||
max_slice_depth,
|
||||
max_memory_slice_depth,
|
||||
self._path_callback,
|
||||
lambda: self.cancelled,
|
||||
)
|
||||
|
||||
@@ -166,7 +166,9 @@ class ConfigView(qtw.QWidget):
|
||||
gen_wid.setLayout(gen_lay)
|
||||
# Find layout
|
||||
fnd_lay = qtw.QGridLayout()
|
||||
for i, name in enumerate(["max_call_level", "max_slice_depth"]):
|
||||
for i, name in enumerate(
|
||||
["max_call_level", "max_slice_depth", "max_memory_slice_depth"]
|
||||
):
|
||||
setting: SpinboxSetting = self.config_ctr.get_setting(name)
|
||||
if not setting:
|
||||
continue
|
||||
|
||||
@@ -84,6 +84,13 @@ class TestData(unittest.TestCase):
|
||||
max_value=9999,
|
||||
help="maximum slice depth to stop the search",
|
||||
),
|
||||
"max_memory_slice_depth": SpinboxSetting(
|
||||
name="max_memory_slice_depth",
|
||||
value=-1,
|
||||
min_value=-1,
|
||||
max_value=9999,
|
||||
help="maximum memory slice depth to stop the search",
|
||||
),
|
||||
"src_highlight_color": ComboboxSetting(
|
||||
name="src_highlight_color",
|
||||
value="Orange",
|
||||
@@ -271,6 +278,7 @@ class TestData(unittest.TestCase):
|
||||
"max_workers",
|
||||
"max_call_level",
|
||||
"max_slice_depth",
|
||||
"max_memory_slice_depth",
|
||||
"max_turns",
|
||||
"max_completion_tokens",
|
||||
]:
|
||||
|
||||
@@ -59,6 +59,7 @@ class TestCase(unittest.TestCase):
|
||||
max_workers: int | None = -1,
|
||||
max_call_level: int = 3,
|
||||
max_slice_depth: int = -1,
|
||||
max_memory_slice_depth: int = -1,
|
||||
enable_all_funs: bool = True,
|
||||
) -> List[Path]:
|
||||
"""
|
||||
@@ -70,6 +71,7 @@ class TestCase(unittest.TestCase):
|
||||
max_workers=max_workers,
|
||||
max_call_level=max_call_level,
|
||||
max_slice_depth=max_slice_depth,
|
||||
max_memory_slice_depth=max_memory_slice_depth,
|
||||
enable_all_funs=enable_all_funs,
|
||||
)
|
||||
slicer.start()
|
||||
@@ -1253,13 +1255,9 @@ class TestMultiThreading(TestCase):
|
||||
bv = bn.load(file)
|
||||
bv.update_analysis_and_wait()
|
||||
# Assert results
|
||||
paths = self.get_paths(
|
||||
bv, max_workers=1, max_call_level=3, enable_all_funs=True
|
||||
)
|
||||
paths = self.get_paths(bv, max_workers=1, enable_all_funs=True)
|
||||
for max_workers in [2, 4, 8, -1]:
|
||||
paths_mt = self.get_paths(
|
||||
bv, max_workers, max_call_level=3, enable_all_funs=True
|
||||
)
|
||||
paths_mt = self.get_paths(bv, max_workers, enable_all_funs=True)
|
||||
self.assertCountEqual(paths, paths_mt, f"{max_workers:d} workers")
|
||||
# Close binary
|
||||
bv.file.close()
|
||||
|
||||
Reference in New Issue
Block a user