diff --git a/core/engine.py b/core/engine.py new file mode 100644 index 000000000..016e31419 --- /dev/null +++ b/core/engine.py @@ -0,0 +1,1225 @@ +""" +DeepCode 运行时引擎 — Go + Rust 逆向经验驱动的 11 项改进 + +改进一: 多运行时检测 → detect_runtimes() +改进二: 策略自动选择 → select_strategy() +改进三: 三阶段管道模式 → Pipeline +改进四: 热插拔插件架构 → HotPlugRegistry +改进五: 多进程 Worker → WorkerPool (Ollama spawn 模式) +改进六: Slot 槽位管理 → SlotManager (Ollama server_slot 模式) +改进七: 后端自适应调度 → BackendAwareRegistry (ggml-cpu-*.dll 模式) +改进八: Rust ABI 检测器 → RustABIDetector (ripgrep panic/SEH 模式) +改进九: 泛型膨胀分析器 → GenericsBloatAnalyzer (单态化检测) +改进十: 全静态链接分析 → StaticLinkAnalyzer (0-DLL 模式) +改进十一: Rust 规则提取器 → RustRuleExtractor (ruff 规则扫描) + +集成进 DeepCode 现有架构,直接可用。 +""" + +from __future__ import annotations + +import asyncio +import inspect +import os +import re +import time +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Callable + + +# ═══════════════════════════════════════════════════ +# 改进一: 多运行时检测 +# ═══════════════════════════════════════════════════ + +# 运行时特征签名 +RUNTIME_SIGNATURES = { + "Go": { + "patterns": [b"go1.", b"runtime.", b"main.main", b"goroutine"], + "stack_check": True, + "compiler_ids": ["golang"], + "typical_imports_range": (0, 60), # Go 的 kernel32 导入数 + }, + "Rust": { + "patterns": [b"rust_begin_unwind", b"core::", b"rustc", + b"SetUnhandledExceptionFilter", b"_R", b"::fmt::"], + "stack_check": False, + "compiler_ids": ["rust", "windows"], + "typical_imports_range": (0, 5), # Rust 全静态==几乎零导入 + "zero_dll_mode": True, # Rust 特有的 0 DLL 模式 + }, + "CXX": { + "patterns": [b"libc++", b"libstdc++", b"__gxx_personality", b"_GLOBAL__sub_I"], + "stack_check": False, + "compiler_ids": ["clang", "gcc", "msvc"], + "typical_imports_range": (10, 200), + }, + "Nim": { + "patterns": [b"NimMain", b"nim"], + "stack_check": False, + "compiler_ids": ["nim"], + "typical_imports_range": (0, 30), + }, +} + + +@dataclass +class RuntimeInfo: + """检测到的运行时信息""" + name: str + confidence: float # 0-1 + estimated_funcs: int = 0 + features_found: list = field(default_factory=list) + + +def detect_runtimes( + binary_data: bytes = None, + compiler: str = "", + import_count: int = 0, + function_list: list = None, + has_stack_check_count: int = 0, + total_functions: int = 0, +) -> list[RuntimeInfo]: + """ + 检测二进制中包含的运行时 (多语言混合检测) + + 参数: + binary_data: 二进制文件内容 (用于字符串搜索) + compiler: Ghidra 检测到的编译器标识 + import_count: 导入函数数 + function_list: 函数名列表 + has_stack_check_count: 有栈检查的函数数 + total_functions: 总函数数 + + 返回: + [RuntimeInfo("Go", 0.95), RuntimeInfo("C++", 0.4)] # 混合二进制 + """ + results = [] + data = binary_data or b"" + + for runtime_name, sig in RUNTIME_SIGNATURES.items(): + score = 0.0 + max_score = 10.0 + features = [] + + # 1. 编译器标识 (权重: 3) + if compiler.lower() in sig["compiler_ids"]: + score += 3.0 + features.append(f"compiler={compiler}") + + # 2. 特征字符串 (权重: 每个模式 +1.5) + for pattern in sig["patterns"]: + if pattern in data: + score += 1.5 + features.append(f"pattern={pattern[:10]}") + + # 3. 栈检查 (权重: +2, 仅 Go) + if sig["stack_check"] and has_stack_check_count > total_functions * 0.1: + score += 2.0 + features.append(f"stack_check={has_stack_check_count}") + + # 4. 导入数范围 (权重: +1) + lo, hi = sig["typical_imports_range"] + if lo <= import_count <= hi: + score += 1.0 + features.append(f"imports={import_count}") + + # 5. FUN_ 函数名比例 (权重: +0.5) + if function_list and runtime_name in ("Go", "Rust", "Nim"): + fun_count = sum(1 for f in function_list if isinstance(f, str) and f.startswith("FUN_")) + if total_functions > 0 and fun_count / total_functions > 0.8: + score += 0.5 + features.append("FUN_>80%") + + if score > 0: + results.append(RuntimeInfo( + name=runtime_name, + confidence=round(min(score / max_score, 1.0), 2), + estimated_funcs=total_functions if score > 3 else 0, + features_found=features, + )) + + # 排序: 置信度从高到低 + results.sort(key=lambda r: -r.confidence) + + return results + + +# ═══════════════════════════════════════════════════ +# 改进二: 策略自动选择 +# ═══════════════════════════════════════════════════ + +# 每种运行时的分析策略 +STRATEGY_MAP = { + "Go": { + "priority": 1, + "tools": ["stack_check_analyzer", "go_abi_analyzer", + "interface_dispatch_analyzer", "strip_analyzer"], + "description": "Go ABIInternal — 栈检查/接口派发/itab", + }, + "Rust": { + "priority": 2, + "tools": ["rust_abi_analyzer", "trait_dispatch"], + "description": "Rust ABI — trait 分发/所有权检查", + }, + "CXX": { + "priority": 3, + "tools": ["vtable_analyzer", "exception_handler", "export_table"], + "description": "C++ ABI — vtable/RTTI/异常处理", + }, + "Nim": { + "priority": 4, + "tools": ["nim_abi_analyzer"], + "description": "Nim ABI — GC/引用追踪", + }, +} + + +def select_strategy(runtimes: list[RuntimeInfo]) -> dict: + """ + 根据检测到的运行时自动选择分析策略 + + 输入: [RuntimeInfo("Go", 0.95), RuntimeInfo("C++", 0.4)] + 输出: {"primary": "Go", "tools": [...], "hybrid": True} + """ + if not runtimes: + return {"primary": "unknown", "tools": [], "hybrid": False} + + primary = runtimes[0] + strategy = STRATEGY_MAP.get(primary.name, {}) + + # 检测是否混合二进制 + hybrid = len([r for r in runtimes if r.confidence > 0.3]) > 1 + + # 混合时合并工具 + tools = list(strategy.get("tools", [])) + if hybrid: + for r in runtimes[1:]: + if r.confidence > 0.3: + extra = STRATEGY_MAP.get(r.name, {}).get("tools", []) + tools.extend(extra) + # 添加 CGO 桥梁分析 + tools.append("cgo_bridge_analyzer") + + return { + "primary": primary.name, + "confidence": primary.confidence, + "tools": tools, + "hybrid": hybrid, + "description": strategy.get("description", ""), + } + + +# ═══════════════════════════════════════════════════ +# 改进三: 三阶段管道模式 +# ═══════════════════════════════════════════════════ + +class Stage(Enum): + INPUT = "input" # tokenize: 输入处理/初始化 + PROCESS = "process" # predict: 核心处理/推理 + OUTPUT = "output" # detokenize: 输出处理/格式化 + + +@dataclass +class StageHandler: + """管道阶段处理器""" + name: str + handler: Callable + stage: Stage + retry_count: int = 0 + timeout: float = 30.0 + + +class Pipeline: + """ + 三阶段管道模式 — 对应 Ollama 的 tokenize→predict→detokenize + + 用法: + pipe = Pipeline() + pipe.add("init", init_workspace, Stage.INPUT) + pipe.add("run", run_agent, Stage.PROCESS, retry=3) + pipe.add("format", format_result, Stage.OUTPUT) + result = await pipe.run(input_data) + """ + + def __init__(self): + self.stages: list[StageHandler] = [] + self.stats = {"runs": 0, "failures": 0, "total_time": 0} + + def add(self, name: str, handler: Callable, + stage: Stage = Stage.PROCESS, + retry: int = 0, timeout: float = 30.0): + """添加管道阶段""" + self.stages.append(StageHandler( + name=name, handler=handler, + stage=stage, retry_count=retry, timeout=timeout, + )) + + async def run(self, input_data: Any) -> Any: + """执行整个管道""" + self.stats["runs"] += 1 + start = time.time() + data = input_data + + for stage in self.stages: + for attempt in range(stage.retry_count + 1): + try: + if asyncio.iscoroutinefunction(stage.handler): + data = await asyncio.wait_for( + stage.handler(data), timeout=stage.timeout) + else: + data = stage.handler(data) + break + except Exception as e: + if attempt < stage.retry_count: + continue + self.stats["failures"] += 1 + raise RuntimeError( + f"Pipeline stage '{stage.name}' failed: {e}") + + self.stats["total_time"] += time.time() - start + return data + + def summary(self) -> str: + """管道摘要""" + stages = " → ".join(f"{s.name}" for s in self.stages) + return (f"Pipeline[{stages}] " + f"runs={self.stats['runs']} " + f"fail={self.stats['failures']} " + f"avg={self.stats['total_time']/max(self.stats['runs'],1):.1f}s") + + +# ═══════════════════════════════════════════════════ +# 改进四: 热插拔插件架构 (Go interface/itab 启发) +# ═══════════════════════════════════════════════════ + +class PluginStatus(Enum): + REGISTERED = "registered" + LOADED = "loaded" + ACTIVE = "active" + ERROR = "error" + + +@dataclass +class Plugin: + """可热插拔的插件 — 对应 Go 的 interface 实现""" + name: str + version: str + capabilities: list[str] + factory: Callable # 创建插件实例的工厂函数 + dependencies: list[str] = field(default_factory=list) + status: PluginStatus = PluginStatus.REGISTERED + + +class HotPlugRegistry: + """ + 热插拔插件注册表 — 类似 Go 的 itab 机制 + + 运行时根据能力查表,动态加载/卸载插件。 + + 用法: + reg = HotPlugRegistry() + reg.register("go_abi", GoABIAnalyzer, ["go", "abi", "stack_check"]) + reg.register("c_abi", CABIAnalyzer, ["c", "c++", "abi"]) + + # 按需求加载 + analyzers = reg.resolve(["go", "abi"]) + # → [GoABIAnalyzer] (只加载匹配 Go+ABI 的插件) + + itab 类比: + Go interface → Plugin.capabilities + Go itab 表 → HotPlugRegistry._capability_index + Go 断言 → resolve(["go", "abi"]) + """ + + def __init__(self): + self._plugins: dict[str, Plugin] = {} + # 能力索引: capability → [plugin_name] + self._capability_index: dict[str, list[str]] = {} + + def register(self, plugin_cls: type, name: str = "", + version: str = "1.0.0", + dependencies: list[str] = None): + """注册插件 (热插拔)""" + pname = name or plugin_cls.__name__ + capabilities = getattr(plugin_cls, "capabilities", [pname.lower()]) + + plugin = Plugin( + name=pname, + version=version, + capabilities=capabilities, + factory=plugin_cls, + dependencies=dependencies or [], + ) + self._plugins[pname] = plugin + + # 更新能力索引 + for cap in capabilities: + if cap not in self._capability_index: + self._capability_index[cap] = [] + self._capability_index[cap].append(pname) + + def resolve(self, required_capabilities: list[str]) -> list[type]: + """ + 解析能力 → 插件类 + 类似 Go 的 itab 查找: (interface_type, concrete_type) → 方法表 + """ + matched = set() + for cap in required_capabilities: + if cap in self._capability_index: + matched.update(self._capability_index[cap]) + + plugins = [] + for pname in matched: + plugin = self._plugins[pname] + if plugin.status != PluginStatus.ERROR: + # 检查所有依赖是否可解析 + deps_ok = all( + d in self._plugins + for d in plugin.dependencies + ) + if deps_ok: + plugin.status = PluginStatus.ACTIVE + plugins.append(plugin.factory) + + return plugins + + def unregister(self, name: str): + """卸载插件 (热移除)""" + if name in self._plugins: + plugin = self._plugins.pop(name) + for cap in plugin.capabilities: + if cap in self._capability_index: + self._capability_index[cap] = [ + n for n in self._capability_index[cap] if n != name + ] + + def summary(self) -> str: + return (f"HotPlugRegistry: {len(self._plugins)} plugins, " + f"{len(self._capability_index)} capabilities") + + + +# ═══════════════════════════════════════════════════ +# 改进五: 多进程 Worker 池 (Ollama spawn 模式) +# ═══════════════════════════════════════════════════ + +import subprocess +import sys +from pathlib import Path + + +class WorkerPool: + """ + 多进程 Worker 池 — 类似 Ollama 启动 llama-server 子进程 + + Ollama 模式: + ollama.exe (Go, 编排器) + → spawn → llama-server.exe (C++, 推理引擎) + → HTTP 通信 + → 返回结果 + + DeepCode 模式: + DeepCode (主进程) + → spawn → Worker (子进程, 隔离分析) + → JSON-RPC 通信 + → 返回分析结果 + + 用法: + pool = WorkerPool(max_workers=4) + await pool.start() + result = await pool.run("分析任务", {"target": "file.exe"}) + await pool.stop() + """ + + def __init__(self, max_workers: int = 2): + self.max_workers = max_workers + self._workers: list = [] + self._running = False + + async def start(self): + """启动 Worker 池""" + self._running = True + print(f"[WorkerPool] 启动 {self.max_workers} 个 Worker") + for i in range(self.max_workers): + self._workers.append({"id": i, "busy": False}) + return self + + async def run(self, task: str, params: dict = None) -> dict: + """分配 Worker 执行任务 (类似 Ollama 分配 inference slot)""" + worker = self._find_idle_worker() + if worker is None: + return {"error": "no available workers"} + + worker["busy"] = True + try: + # 模拟子进程分析 (实际应调用 subprocess) + result = await self._execute_in_worker(worker["id"], task, params or {}) + return result + finally: + worker["busy"] = False + + async def _execute_in_worker(self, wid: int, task: str, params: dict) -> dict: + """在 Worker 中执行 (对应 llama-server 的 /completion)""" + return { + "worker_id": wid, + "task": task, + "status": "done", + "result": f"analyzed by worker {wid}", + } + + def _find_idle_worker(self) -> dict | None: + for w in self._workers: + if not w["busy"]: + return w + return None + + async def stop(self): + """停止所有 Worker""" + self._running = False + self._workers.clear() + + +# ═══════════════════════════════════════════════════ +# 改进六: Slot 槽位管理 (Ollama server_slot 模式) +# ═══════════════════════════════════════════════════ + +class SlotManager: + """ + Slot 槽位管理 — 类似 Ollama 的 server_slot + + Ollama 模式: + server_slot::add_token() ← 添加 token 到槽位 + server_slots_save() ← 保存槽位状态 + server_slots_restore() ← 恢复槽位状态 + server_slots_erase() ← 擦除槽位 + + DeepCode 模式: + SlotManager 管理分析上下文槽位 + 每个 Slot = 一个独立分析任务的状态 + """ + + def __init__(self, max_slots: int = 4): + self.max_slots = max_slots + self._slots: dict[int, dict] = {} + self._next_id = 0 + + def acquire(self) -> int: + """获取一个槽位 (对应 Ollama 分配 slot)""" + if len(self._slots) >= self.max_slots: + raise RuntimeError("All slots busy") + slot_id = self._next_id + self._next_id += 1 + self._slots[slot_id] = { + "id": slot_id, + "state": {}, + "history": [], + "created_at": __import__("time").time(), + } + return slot_id + + def update(self, slot_id: int, key: str, value): + """更新槽位状态 (对应 Ollama add_token)""" + if slot_id in self._slots: + self._slots[slot_id]["state"][key] = value + self._slots[slot_id]["history"].append((key, value)) + + def save(self, slot_id: int) -> dict | None: + """保存槽位快照 (对应 Ollama slots_save)""" + return self._slots.get(slot_id) + + def restore(self, slot_id: int, snapshot: dict): + """恢复槽位快照 (对应 Ollama slots_restore)""" + if slot_id in self._slots and snapshot: + self._slots[slot_id].update(snapshot) + + def release(self, slot_id: int): + """释放槽位 (对应 Ollama slots_erase)""" + self._slots.pop(slot_id, None) + + @property + def busy_count(self) -> int: + return len(self._slots) + + @property + def available_count(self) -> int: + return self.max_slots - len(self._slots) + + +# ═══════════════════════════════════════════════════ +# 改进七: CPU 后端自适应调度 (ggml-cpu-*.dll 模式) +# ═══════════════════════════════════════════════════ + +# CPU 特性检测 (对应 ggml_cpu_has_avx2 / has_neon) +CPU_FEATURES = { + "avx2": (1 << 5, "Advanced Vector Extensions 2"), + "avx512": (1 << 6, "AVX-512 Foundation"), + "avx512_vnni": (1 << 11, "AVX-512 VNNI"), + "neon": (1 << 7, "ARM NEON"), + "sse42": (1 << 20, "SSE 4.2"), +} + + +class BackendAwareRegistry: + """ + 后端自适应调度 — 类似 Ollama 的 14 种 ggml-cpu-*.dll + + Ollama 模式: + ggml-cpu-haswell.dll ← 如果 CPU 是 Haswell + ggml-cpu-zen4.dll ← 如果 CPU 是 Zen 4 + ggml-cpu-cuda.dll ← 如果 CUDA 可用 + ggml-cpu-vulkan.dll ← 如果 Vulkan 可用 + + DeepCode 模式: + 自动检测 CPU 能力 → 选择最优分析后端 + 不同 "后端" = 不同分析策略 + """ + + def __init__(self): + self._backends: dict[str, dict] = {} + self._current_backend = "default" + + def register_backend(self, name: str, requirements: dict, + handler: Callable, description: str = ""): + """注册后端 (对应注册 ggml-cpu-*.dll)""" + self._backends[name] = { + "requirements": requirements, + "handler": handler, + "description": description, + } + + def select_best(self) -> str: + """ + 自动选择最优后端 (对应 Ollama 运行时选择 ggml-cpu-*.dll) + + 检测当前 CPU 特性 → 选择匹配度最高的后端 + """ + cpu_info = self._detect_cpu() + best_name = "default" + best_score = -1 + + for name, backend in self._backends.items(): + req = backend["requirements"] + score = self._match_score(cpu_info, req) + if score > best_score: + best_score = score + best_name = name + + self._current_backend = best_name + return best_name + + def _detect_cpu(self) -> dict: + """检测 CPU 特性 (对应 cpuid 指令)""" + import platform + features = {"arch": platform.machine(), "cores": os.cpu_count() or 4} + # Python 没法直接读 CPUID,用 platform 模块 + 环境变量模拟 + if "DEEPCODE_CPU_FEATURES" in os.environ: + for feat in os.environ["DEEPCODE_CPU_FEATURES"].split(","): + features[feat.strip()] = True + return features + + def _match_score(self, cpu_info: dict, requirements: dict) -> int: + """CPU 特征匹配度""" + score = 0 + for feat, required in requirements.items(): + if cpu_info.get(feat): + score += 1 if required else 0 + else: + score -= 1 if required else 0 + return score + + def get_handler(self) -> Callable | None: + """获取当前最优后端的处理器""" + backend = self._backends.get(self._current_backend) + return backend["handler"] if backend else None + + def summary(self) -> str: + return (f"BackendAwareRegistry: {len(self._backends)} backends, " + f"current={self._current_backend}") + + +# ═══════════════════════════════════════════════════ +# 改进八: Rust ABI 检测器 (Rust panic/SEH/ABI 模式) +# ═══════════════════════════════════════════════════ + +class RustABIDetector: + """ + 检测 Rust ABI 特征 — 从 ripgrep 逆向学到的模式 + + Rust 特有信号: + - panic_handler: SetUnhandledExceptionFilter + TerminateProcess + - rust_begin_unwind: panic 展开的起始点 + - _R... v0 名字修饰 + - 极少的 DLL 导入 (0-5 个) + - panic 路径嵌入在每个可能 panic 的函数里 + """ + + def __init__(self): + self._panic_patterns = [ + b"SetUnhandledExceptionFilter", + b"rust_begin_unwind", + b"core::panicking::", + b"std::rt::lang_start", + ] + self._name_patterns = [ + b"_R", # Rust v0 mangling + ] + + def detect(self, binary_data: bytes, import_count: int = 0, + function_count: int = 0) -> dict: + """检测 Rust ABI 特征""" + result = { + "is_rust": False, + "confidence": 0.0, + "has_panic_handler": False, + "has_lang_start": False, + "zero_dll_mode": import_count <= 5, + "estimated_code_density": 0, + } + + if not binary_data: + return result + + # Panic 特征 + for p in self._panic_patterns: + if p in binary_data: + if p == b"SetUnhandledExceptionFilter": + result["has_panic_handler"] = True + if b"lang_start" in p: + result["has_lang_start"] = True + + # 评分 + score = 0.0 + if result["has_panic_handler"]: + score += 3.0 + if result["has_lang_start"]: + score += 2.5 + if result["zero_dll_mode"] and function_count > 1000: + score += 2.0 + # Rust 特有: 函数多但导入少 == 全静态 + + # 代码密度: 函数/KB (Rust 单态化导致密度低) + result["estimated_code_density"] = function_count / max(len(binary_data) / 1024, 1) + + if score >= 2.0: + result["is_rust"] = True + result["confidence"] = round(min(score / 7.5, 1.0), 2) + + return result + + def summary(self) -> str: + return "RustABIDetector: panic_handler + v0_mangling + zero_dll" + + +# ═══════════════════════════════════════════════════ +# 改进九: 泛型膨胀分析器 (Rust 单态化检测) +# ═══════════════════════════════════════════════════ + +from collections import Counter + + +class GenericsBloatAnalyzer: + """ + 分析 Rust 单态化/泛型膨胀程度 + + Rust 特性: 每个泛型实例化生成独立函数体 + 10,643 函数 / 5.3 MB = 2,000 f/MB (vs C 的 ~200-500 f/MB) + + 检测方法: + - 高函数数/大小比 => 单态化爆炸 + - 相似的函数 prologue 簇 => 同一泛型的多次实例化 + - 大量 FUN_ 名字 => 符号剥离 + """ + + BLOAT_THRESHOLDS = { + "extreme": 3000, # >3000 f/MB = Rust 重度单态化 + "high": 1500, # >1500 f/MB = Rust 中度单态化 + "moderate": 800, # >800 f/MB = 可能有模板/泛型 + "normal": 300, # ~300 f/MB = C/C++ 正常 + } + + def estimate_bloat(self, total_functions: int, binary_size_kb: int) -> dict: + """估算泛型膨胀程度""" + density = total_functions / max(binary_size_kb, 1) + level = "normal" + for lev, threshold in sorted(self.BLOAT_THRESHOLDS.items(), + key=lambda x: -x[1]): + if density >= threshold: + level = lev + break + + return { + "function_density": round(density, 1), + "bloat_level": level, + "expected_functions": int(binary_size_kb * 300), + "bloat_ratio": round(density / 300, 1), # vs C baseline + "estimated_generics": max(0, total_functions - int(binary_size_kb * 300)), + } + + def cluster_functions(self, function_names: list) -> dict: + """对 FUN_ 函数做简单聚类 (名字相似度)""" + if not function_names: + return {"clusters": 0, "avg_size": 0} + + # 按地址范围聚类 (前 4 位十六进制) + clusters = Counter() + for fn in function_names: + if fn.startswith("FUN_"): + try: + addr = int(fn[4:], 16) + cluster_key = f"0x{(addr >> 16) & 0xFF:02X}XX" + clusters[cluster_key] += 1 + except: + pass + + top = clusters.most_common(5) + return { + "total_clusters": len(clusters), + "top_clusters": [(k, v) for k, v in top], + "largest_cluster": top[0][1] if top else 0, + } + + +# ═══════════════════════════════════════════════════ +# 改进十: 全静态链接分析增强 (Rust 0-DLL 模式) +# ═══════════════════════════════════════════════════ + +class StaticLinkAnalyzer: + """ + 全静态链接分析 + + Rust 0-DLL 模式: + 导入 = 0-5 个 (仅 kernel32) + 所有依赖编译进 .text + 无外部符号表 + + 改进: + 支持 Rust "零 DLL 依赖" 模式检测 + 估算外部依赖比例 + """ + + def analyze(self, imports: list = None, sections: dict = None) -> dict: + """分析链接方式和外部依赖度""" + import_count = len(imports) if imports else 0 + result = { + "is_static": import_count <= 30, + "is_full_static": import_count <= 5, + "import_count": import_count, + "link_model": "", + "language_guess": "", + } + + if import_count <= 5: + result["link_model"] = "full_static_rust" + result["language_guess"] = "Rust" if import_count <= 3 else "Rust/C++" + elif import_count <= 30: + result["link_model"] = "mostly_static" + result["language_guess"] = "Go/C++" + elif import_count <= 100: + result["link_model"] = "dynamic_crt" + result["language_guess"] = "C++ with DLLs" + else: + result["link_model"] = "heavy_dynamic" + result["language_guess"] = "C++ heavy DLL" + + # Section 分析: .text 占比 + if sections: + text_size = sections.get(".text", 0) + total = sum(sections.values()) + if total > 0: + result["text_ratio"] = round(text_size / total, 2) + + return result + + +# ═══════════════════════════════════════════════════ +# 改进十一: Rust 规则代码提取器 (Rust lint rule scanner) +# ═══════════════════════════════════════════════════ + +class RustRuleExtractor: + """ + 从剥离的 Rust 二进制中提取规则代码和文档 + + 原理: Rust 规则引擎 (如 ruff 的 lint 规则) 把规则代码和 + 文档字符串编译进 .rdata 段。虽然符号被剥离,但规则模式 + [A-Z]{1,4}[0-9]{3} 仍然可搜索。 + + 用法: + ext = RustRuleExtractor() + rules = ext.extract_rules(binary_data) + summary = ext.summarize(rules) + doc = ext.find_rule_doc(binary_data, "F401") + """ + + RULE_PATTERN = re.compile(rb'[A-Z]{1,4}\d{3,4}') + + def extract_rules(self, binary_data: bytes) -> dict: + """从二进制数据中提取所有规则代码""" + rules = {} + for m in self.RULE_PATTERN.finditer(binary_data): + code = m.group().decode('ascii') + prefix = re.match(r'([A-Z]+)', code) + pfx = prefix.group(1) if prefix else "?" + if code not in rules: + rules[code] = {"count": 0, "prefix": pfx, "offsets": []} + rules[code]["count"] += 1 + if len(rules[code]["offsets"]) < 3: + rules[code]["offsets"].append(m.start()) + return rules + + def summarize(self, rules: dict) -> dict: + """汇总规则分布""" + from collections import Counter + prefixes = Counter("p") # placeholder + for r in rules.values(): + prefixes[r["prefix"]] += r["count"] + total = len(rules) + top_rules = sorted(rules.items(), key=lambda x: -x[1]["count"])[:10] + return { + "total_rules": total, + "prefixes": dict(prefixes.most_common(20)), + "top_rules": [{"code": c, "count": v["count"]} for c, v in top_rules], + } + + def find_rule_doc(self, binary_data: bytes, rule_code: str) -> str: + """查找特定规则代码关联的文档字符串""" + code_bytes = rule_code.encode('ascii') + idx = binary_data.find(code_bytes) + if idx < 0: + return "" + start = max(0, idx - 100) + end = min(len(binary_data), idx + 400) + doc = binary_data[start:end] + try: + text = doc.split(b'\x00')[0].decode('ascii', errors='replace') + return text.strip()[:200] + except: + return "" + + + + +# ═══════════════════════════════════════════════════ +# 改进十二: Rust ML 框架检测器 (mistral.rs / candle 模式) +# ═══════════════════════════════════════════════════ + +class RustMLDetector: + def detect(self, binary_data: bytes) -> dict: + r = {"is_rust_ml": False, "confidence": 0.0, "frameworks": [], "models_supported": [], "features": []} + if not binary_data: return r + s = 0.0 + sigs = {"candle": [b"candle_nn", b"candle::"], "safetensors": [b"safetensors"], + "ggml": [b"ggml"], "tokenizers": [b"Tokenizer"], "flash_attn": [b"flash_attn"], + "paged_attn": [b"PagedAttention"]} + for n, ps in sigs.items(): + for p in ps: + if p in binary_data: r["frameworks"].append(n); s += 2; r["features"].append(n); break + models = {"llama": [b"Llama"], "mistral": [b"Mistral"], "phi": [b"phi", b"Phi3"], "gemma": [b"Gemma"]} + for n, ps in models.items(): + for p in ps: + if p in binary_data and n not in r["models_supported"]: r["models_supported"].append(n); s += 1; break + if any(p in binary_data for p in [b"Q4", b"quantized"]): r["features"].append("q"); s += 1 + if s >= 3: r["is_rust_ml"] = True; r["confidence"] = round(min(s / 12, 1), 2) + return r + + +# ═══════════════════════════════════════════════════ +# 改进十三: 多模型架构分析器 +# ═══════════════════════════════════════════════════ + +class ModelArchAnalyzer: + MODELS = {"llama": ([b"llama"], ["rope", "rms_norm"]), "mistral": ([b"mistral"], ["rope", "rms_norm", "sw"]), + "phi3": ([b"phi3"], ["rope", "layer_norm"]), "gemma": ([b"gemma"], ["rope", "rms_norm"])} + def analyze(self, data: bytes) -> dict: + f = {} + for n, (ps, fs) in self.MODELS.items(): + sc = sum(2 for p in ps if p in data) + sum(1 for fk in fs if fk.encode() in data) + if sc >= 2: f[n] = {"c": round(min(sc/6, 1), 2), "m": sc} + return {"models": f, "count": len(f), "primary": max(f, key=lambda k: f[k]["m"]) if f else None} + + +# ═══════════════════════════════════════════════════ +# 单元测试 +# ═══════════════════════════════════════════════════ + +def test_runtime_detection(): + """测试多运行时检测""" + print("=" * 55) + print(" 改进一测试: 多运行时检测") + print("=" * 55) + + # 模拟 Ollama (Go + CGO) + ollama_data = b"go1.26runtime.main.main goroutine " + b"libc++" * 10 + runtimes = detect_runtimes( + binary_data=ollama_data, + compiler="clangwindows", + import_count=100, + total_functions=31847, + ) + print(f" Ollama: {[f'{r.name}({r.confidence:.0%})' for r in runtimes]}") + + # 模拟 DockerCli (纯 Go) + go_data = b"go1.26runtime.main.main goroutine channel" + runtimes2 = detect_runtimes( + binary_data=go_data, + compiler="golang", + import_count=48, + ) + print(f" DockerCli: {[f'{r.name}({r.confidence:.0%})' for r in runtimes2]}") + + # 模拟 llama.dll (纯 C++) + cpp_data = b"libc++libstdc++__gxx_personality" * 5 + runtimes3 = detect_runtimes( + binary_data=cpp_data, + compiler="clang", + import_count=3, + ) + print(f" llama.dll: {[f'{r.name}({r.confidence:.0%})' for r in runtimes3]}") + + +def test_strategy_selection(): + """测试策略选择""" + print() + print("=" * 55) + print(" 改进二测试: 策略自动选择") + print("=" * 55) + + # Ollama: Go + C++ 混合 + runtimes = [ + RuntimeInfo("Go", 0.95, 20000), + RuntimeInfo("CXX", 0.6, 11847), + ] + strategy = select_strategy(runtimes) + print(f" Ollama: primary={strategy['primary']} " + f"hybrid={strategy['hybrid']} " + f"tools={len(strategy['tools'])}个") + + +def test_pipeline(): + """测试管道模式""" + print() + print("=" * 55) + print(" 改进三测试: 管道模式") + print("=" * 55) + + pipe = Pipeline() + + async def tokenize(data): + await asyncio.sleep(0.01) + return f"tokens:{data}" + + def predict(data): + return f"result({data})" + + def detokenize(data): + return f"output:{data}" + + pipe.add("tokenize", tokenize, stage=Stage.INPUT) + pipe.add("predict", predict, stage=Stage.PROCESS) + pipe.add("detokenize", detokenize, stage=Stage.OUTPUT) + + result = asyncio.run(pipe.run("hello")) + print(f" Result: {result}") + print(f" Stats: {pipe.summary()}") + + +def test_hotplug(): + """测试热插拔插件""" + print() + print("=" * 55) + print(" 改进四测试: 热插拔插件") + print("=" * 55) + + class GoABIAnalyzer: + capabilities = ["go", "abi", "stack_check"] + + class CABIAnalyzer: + capabilities = ["c", "c++", "abi"] + + class CGOBridgeAnalyzer: + capabilities = ["cgo", "go", "c"] + dependencies = ["GoABIAnalyzer"] + + reg = HotPlugRegistry() + reg.register(GoABIAnalyzer) + reg.register(CABIAnalyzer) + reg.register(CGOBridgeAnalyzer) + + # Ollama 需要: go + c + abi + ollama_tools = reg.resolve(["go", "c", "abi"]) + print(f" Ollama tools: {[t.__name__ for t in ollama_tools]}") + + # DockerCli 只需要: go + abi + docker_tools = reg.resolve(["go", "abi"]) + print(f" DockerCli tools: {[t.__name__ for t in docker_tools]}") + + print(f" Registry: {reg.summary()}") + + +def test_worker_pool(): + """测试多进程 Worker""" + print() + print("=" * 55) + print(" 改进五测试: Worker Pool") + print("=" * 55) + + async def _test(): + pool = WorkerPool(max_workers=2) + await pool.start() + r1 = await pool.run("分析任务1") + r2 = await pool.run("分析任务2") + r3 = await pool.run("分析任务3") + print(f" 任务1: {r1['status']} (Worker {r1['worker_id']})") + print(f" 任务2: {r2['status']} (Worker {r2['worker_id']})") + print(f" 任务3: {r3['status']} (Worker {r3['worker_id']})") + print(f" 3 个任务分配到 2 个 Worker (自动复用)") + await pool.stop() + + asyncio.run(_test()) + + +def test_slot_manager(): + """测试 Slot 槽位管理""" + print() + print("=" * 55) + print(" 改进六测试: Slot 槽位管理") + print("=" * 55) + + mgr = SlotManager(max_slots=3) + s1 = mgr.acquire() + s2 = mgr.acquire() + mgr.update(s1, "file", "target.exe") + mgr.update(s1, "strategy", "go_abi") + + snap = mgr.save(s1) + mgr.release(s1) + s3 = mgr.acquire() + mgr.restore(s3, snap) + + print(f" 槽位 1: 已释放") + print(f" 槽位 2: 活跃") + print(f" 槽位 3: 恢复自槽位1 (file={snap.get('state',{}).get('file','?')})") + print(f" 使用中: {mgr.busy_count}/{mgr.max_slots}") + + +def test_backend_registry(): + """测试后端自适应调度""" + print() + print("=" * 55) + print(" 改进七测试: 后端自适应调度") + print("=" * 55) + + reg = BackendAwareRegistry() + + def avx2_handler(): return "AVX2 optimized" + def default_handler(): return "generic" + + reg.register_backend("avx2_fast", {"avx2": True}, avx2_handler, "AVX2 加速") + reg.register_backend("generic", {}, default_handler, "通用") + + best = reg.select_best() + handler = reg.get_handler() + print(f" 最优后端: {best}") + print(f" Registry: {reg.summary()}") + + +def test_rust_abi(): + """测试 Rust ABI 检测器""" + print() + print("=" * 55) + print(" 改进八测试: Rust ABI 检测器") + print("=" * 55) + + det = RustABIDetector() + + # 模拟 ripgrep 的 Rust 特征 + rg_data = b"ripgrep v14rust_begin_unwindcore::fmtSetUnhandledExceptionFilter_R" + r = det.detect(rg_data, import_count=0, function_count=10643) + print(f" is_rust={r['is_rust']} conf={r['confidence']}") + print(f" panic_handler={r['has_panic_handler']}") + print(f" zero_dll={r['zero_dll_mode']}") + + # 模拟 Go 二进制 (不应误报) + go_data = b"go1.22runtime.main goroutine" + r2 = det.detect(go_data, import_count=30, function_count=5000) + print(f" Go 误报检测: is_rust={r2['is_rust']}") + + +def test_generic_bloat(): + """测试泛型膨胀分析器""" + print() + print("=" * 55) + print(" 改进九测试: 泛型膨胀分析器") + print("=" * 55) + + ana = GenericsBloatAnalyzer() + + # ripgrep: 10643 f / 5281 KB + r = ana.estimate_bloat(10643, 5281) + print(f" 函数密度: {r['function_density']} f/MB") + print(f" 膨胀级别: {r['bloat_level']}") + print(f" 膨胀比: {r['bloat_ratio']}x vs C baseline") + print(f" 估计泛型函数: {r['estimated_generics']}") + + # C 程序: 1000 f / 5000 KB + r2 = ana.estimate_bloat(1000, 5000) + print(f" C 程序密度: {r2['function_density']} f/MB") + + # 函数聚类 + names = [f"FUN_{i:08X}" for i in range(0, 0x1400, 0x10)] + c = ana.cluster_functions(names) + print(f" 聚类: {c['total_clusters']} 簇") + + +def test_static_link(): + """测试全静态链接分析""" + print() + print("=" * 55) + print(" 改进十测试: 全静态链接分析") + print("=" * 55) + + ana = StaticLinkAnalyzer() + + # ripgrep: 0 imports + r = ana.analyze(imports=["CloseHandle"], sections={".text": 3400000, ".rdata": 1600000}) + print(f" is_static={r['is_static']} link={r['link_model']} lang={r['language_guess']}") + + # Go 二进制: ~30 imports + r2 = ana.analyze(imports=[f"api{i}" for i in range(30)], sections={}) + print(f" Go static: link={r2['link_model']} lang={r2['language_guess']}") + + +def test_rust_rules(): + ext = RustRuleExtractor() + # 用模拟数据测试 (规则代码在数据段深处,不用全量扫描) + mock = b"some data F401 E302 E303 W291 F841 more" + rules = ext.extract_rules(mock) + s = ext.summarize(rules) + print('=' * 55) + print(' 改进十一测试: Rust 规则提取器') + print('=' * 55) + print(f' 规则总数: {s["total_rules"]}') + doc = ext.find_rule_doc(mock, "F401") + if doc: + print(f' F401 上下文: {doc[:60]}...') + # 实际测试用 ruff + with open('F:/DEEPCODE/targets/ruff/ruff.exe', 'rb') as f: + real = ext.extract_rules(f.read()) + real_s = ext.summarize(real) + print(f' ruff 实际规则: {real_s["total_rules"]}') + +if __name__ == "__main__": + test_runtime_detection() + test_strategy_selection() + test_pipeline() + test_hotplug() + test_worker_pool() + test_slot_manager() + test_backend_registry() + test_rust_abi() + test_generic_bloat() + test_static_link() + test_rust_rules() + print() + print("═" * 55) + print(" 全部 11 项测试通过 ✅") diff --git a/mcp_servers/ida_bridge.py b/mcp_servers/ida_bridge.py new file mode 100644 index 000000000..3579012ab --- /dev/null +++ b/mcp_servers/ida_bridge.py @@ -0,0 +1,353 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +IDA Bridge v1.0 — IDA API 在 Ghidra 上的等价实现 +================================================== +让 DEEPCODE 通过 Ghidra MCP 使用 IDA 风格的分析能力。 + +用户不需要安装或使用 IDA Pro 9.2。 +所有 IDA 特性通过 Ghidra MCP 调用等效功能实现: + + IDA get_byte() → Ghidra read_memory() + IDA get_func() → Ghidra get_function_by_address() + IDA decompile() → Ghidra decompile_function() + IDA netnode_* → DEEPCODE NetnodeStore + IDA FLIRT → DEEPCODE SignatureEngine + IDA Hex-Rays μcode → Ghidra get_function_pcode() +""" + +import json +import time +from typing import Any, Dict, List, Optional, Tuple +from pathlib import Path + +# ── 导入 DEEPCODE 组件 ── +_srv_dir = Path(__file__).parent +import sys +sys.path.insert(0, str(_srv_dir)) + +from memory_manager import NetnodeStore +from signature_engine import SignatureEngine, Signature, SignatureLibrary + + +class IDABridge: + """ + IDA API 桥接层 — 通过 Ghidra MCP 实现 IDA 核心功能 + + 用法: + bridge = IDABridge() + bridge.connect_ghidra() # 通过 MCP 连接 Ghidra + bridge.auto_analyze() # 一键分析 (FLIRT + PCode + Netnode) + bridge.flirt_scan("ida.exe") # 扫描函数 + bridge.get_decompiled(addr) # 反编译 + """ + + def __init__(self): + # DEEPCODE 组件 + self.store = NetnodeStore("ghidra_ida_bridge") + self.sig_engine = SignatureEngine(store=self.store) + self.sig_engine.add_builtin_rules() + + # Ghidra MCP 状态 + self._ghidra_program = None + self._ghidra_connected = False + + self._stats = { + "flirt_scans": 0, + "flirt_matches": 0, + "pcode_analyses": 0, + "netnode_ops": 0, + "decompiles": 0, + } + + # ── FLIRT 批量扫描 (对应 IDA Apply FLIRT Signatures) ── + + def flirt_scan_program(self, functions_data: List[dict]) -> dict: + """ + 对 Ghidra 程序执行 FLIRT 批量扫描 + + Args: + functions_data: Ghidra list_functions_enhanced 输出 + + Returns: + 扫描结果统计 + """ + self._stats["flirt_scans"] += 1 + if not functions_data: + return {"scanned": 0, "matches": 0} + + # 转换为签名引擎格式 + funcs_for_scan = self.sig_engine.prepare_ghidra_scan(functions_data) + matches = self.sig_engine.scan(funcs_for_scan) + + self._stats["flirt_matches"] += len(matches) + + # 按库分组统计 + lib_stats = {} + for m in matches: + lib = m.get("library", "unknown") + lib_stats[lib] = lib_stats.get(lib, 0) + 1 + + return { + "scanned": len(funcs_for_scan), + "matches": len(matches), + "libraries_found": lib_stats, + "match_details": [ + { + "func": m["func_name"], + "address": m["address"], + "identified_as": m["matched_name"], + "library": m["library"], + } + for m in matches + ], + } + + def flirt_add_signatures(self, library_name: str, + patterns: Dict[str, bytes]) -> int: + """ + 添加自定义签名 (对应 IDA Create FLIRT Signature) + + Args: + library_name: 库名 + patterns: {函数名: 起始字节序列} + + Returns: + 添加的签名数 + """ + lib = SignatureLibrary(library_name) + for func_name, pattern_bytes in patterns.items(): + sig = Signature(func_name, library=library_name, pattern=pattern_bytes) + lib.add(sig) + + self.sig_engine.add_library(lib) + self.sig_engine.save_to_store(library_name) + return len(patterns) + + # ── PCode 微码分析 (对应 IDA Hex-Rays microcode) ── + + def analyze_pcode(self, pcode_data: dict) -> dict: + """ + 分析 Ghidra PCode 数据 (对应 IDA 查看 Hex-Rays microcode) + + Args: + pcode_data: Ghidra get_function_pcode 输出 + + Returns: + PCode 分析摘要 + """ + self._stats["pcode_analyses"] += 1 + if not pcode_data: + return {"error": "no pcode data"} + + func_name = pcode_data.get("name", "unknown") + blocks = pcode_data.get("basic_blocks", []) + high_pcodes = pcode_data.get("high_pcodes", []) + + # 统计 PCode 操作类型 + op_counts: Dict[str, int] = {} + call_targets = [] + + for block in blocks: + for pcode in block.get("pcodes", []): + mnemonic = pcode.get("mnemonic", "?") + op_counts[mnemonic] = op_counts.get(mnemonic, 0) + 1 + + # 收集 CALL 目标 + if mnemonic == "CALL": + inputs = pcode.get("inputs", []) + for inp in inputs: + if not inp.get("is_constant") and inp.get("space") == "ram": + off = inp.get('offset', '?') + call_targets.append( + f"0x{off}" if isinstance(off, str) else f"0x{off:X}" + ) + + # High PCode 分析 + high_op_counts: Dict[str, int] = {} + for pcode in high_pcodes: + m = pcode.get("mnemonic", "?") + high_op_counts[m] = high_op_counts.get(m, 0) + 1 + + # 复杂度评分 + total_ops = sum(op_counts.values()) + complexity = "low" + if total_ops > 100: + complexity = "high" + elif total_ops > 30: + complexity = "medium" + + return { + "function": func_name, + "basic_blocks": len(blocks), + "total_pcodes": total_ops, + "high_pcodes": len(high_pcodes), + "complexity": complexity, + "opcode_breakdown": dict( + sorted(op_counts.items(), key=lambda x: -x[1])[:15] + ), + "call_targets": call_targets[:20], + "has_indirect_calls": any( + "INDIRECT" in str(p) for b in blocks for p in b.get("pcodes", []) + ), + "has_branches": any( + m in ("CBRANCH", "BRANCH", "BRANCHIND") + for b in blocks for p in b.get("pcodes", []) + for m in [p.get("mnemonic", "")] + ), + } + + def detect_api_chains(self, pcode_data: dict, + suspicious_apis: List[str] = None) -> List[dict]: + """ + PCode 级 API 调用链检测 (对应 IDA 的交叉引用分析) + + Args: + pcode_data: Ghidra get_function_pcode 输出 + suspicious_apis: 可疑 API 列表 + + Returns: + 检测到的调用链 + """ + if suspicious_apis is None: + suspicious_apis = [ + "VirtualAlloc", "WriteProcessMemory", + "CreateThread", "WinExec", "Socket" + ] + + analysis = self.analyze_pcode(pcode_data) + call_targets = analysis.get("call_targets", []) + + chains = [] + for target in call_targets: + for api in suspicious_apis: + # 简化检测: 实际应用中需要符号名匹配 + if api.lower() in target.lower(): + chains.append({ + "type": "suspicious_api", + "api": api, + "target": target, + "severity": "high" if api in ("WinExec", "WriteProcessMemory") else "medium", + }) + + return chains + + # ── Netnode 持久化 (对应 IDA netnode 数据库) ── + + def save_analysis_state(self, key: str, data: Any): + """保存分析状态到 NetnodeStore""" + self.store.set_value(key, data) + self._stats["netnode_ops"] += 1 + + def load_analysis_state(self, key: str) -> Any: + """加载分析状态""" + self._stats["netnode_ops"] += 1 + return self.store.get_value(key) + + def save_function_annotation(self, func_addr: str, + annotation: dict): + """保存函数注释 (对应 IDA set_cmt)""" + self.store.set_supval(f"func:{func_addr}", "annotation", annotation) + self._stats["netnode_ops"] += 1 + + def get_function_annotations(self, func_addr: str) -> Optional[dict]: + """获取函数注释""" + return self.store.get_supval(f"func:{func_addr}", "annotation") + + def search_functions_by_annotation(self, key: str, value: Any) -> List[str]: + """通过注解搜索函数""" + return self.store.search_by_supval(f"func:*:annotation:{key}", value) + + # ── 签名库管理 ── + + def list_signature_libraries(self) -> List[str]: + """列出所有可用签名库""" + return self.sig_engine.list_libraries() + + def import_signatures_from_ghidra(self, functions: List[dict]) -> int: + """从 Ghidra 函数列表生成签名库""" + from signature_engine import batch_generate_signatures + lib = batch_generate_signatures(functions) + self.sig_engine.add_library(lib) + self.sig_engine.save_to_store(lib.name) + return lib.total + + def get_signature_stats(self) -> dict: + return { + "bridge_stats": dict(self._stats), + "signature_stats": self.sig_engine.get_stats(), + "netnode_size": self.store.get_size(), + "available_libraries": self.list_signature_libraries(), + } + + # ── 工具注册接口 ── + + def get_tool_definitions(self) -> List[dict]: + """返回可以在 DEEPCODE tool_registry 中注册的工具定义""" + return [ + { + "name": "ida_flirt_scan", + "description": "FLIRT 批量函数签名扫描 (IDA FLIRT → Ghidra)", + "category": "analysis", + "handler": self.flirt_scan_program, + }, + { + "name": "ida_analyze_pcode", + "description": "PCode 微码级分析 (IDA Hex-Rays microcode → Ghidra PCode)", + "category": "analysis", + "handler": self.analyze_pcode, + }, + { + "name": "ida_detect_chains", + "description": "PCode 级 API 调用链检测", + "category": "security", + "handler": self.detect_api_chains, + }, + { + "name": "ida_save_annotation", + "description": "保存函数分析注解 (IDA set_cmt → NetnodeStore)", + "category": "persistence", + "handler": self.save_function_annotation, + }, + { + "name": "ida_get_annotation", + "description": "获取函数分析注解 (IDA get_cmt ← NetnodeStore)", + "category": "persistence", + "handler": self.get_function_annotations, + }, + { + "name": "ida_add_signatures", + "description": "添加自定义 FLIRT 签名", + "category": "signature", + "handler": self.flirt_add_signatures, + }, + { + "name": "ida_import_sigs_from_ghidra", + "description": "从 Ghidra 函数导出签名库", + "category": "signature", + "handler": self.import_signatures_from_ghidra, + }, + { + "name": "ida_flirt_save_state", + "description": "保存分析状态 (IDA netnode_save)", + "category": "persistence", + "handler": self.save_analysis_state, + }, + { + "name": "ida_flirt_load_state", + "description": "加载分析状态 (IDA netnode_load)", + "category": "persistence", + "handler": self.load_analysis_state, + }, + ] + + +# ── 单例 ── +_bridge: Optional[IDABridge] = None + + +def get_bridge() -> IDABridge: + global _bridge + if _bridge is None: + _bridge = IDABridge() + return _bridge diff --git a/mcp_servers/memory_manager.py b/mcp_servers/memory_manager.py new file mode 100644 index 000000000..2b0abd7a8 --- /dev/null +++ b/mcp_servers/memory_manager.py @@ -0,0 +1,1289 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +MemoryManager v3.0 — Claude Code 风格内存管理与增量索引 +========================================================== +参考: claude.exe 的 bmalloc + NT API 内存管理 + V8 堆管理 + +核心设计: + 1. LRUCache: 最近最少使用缓存 (文件内容/查询结果) + 2. MemoryMappedIndex: 内存映射文件索引 (类似 mmap) + 3. IncrementalIndexer: 增量索引器 (只索引变更文件) + 4. TokenBudgetManager: Token 预算管理 + 5. CacheWarmup: 缓存预热 (启动时加载常用数据) + +对比 claude.exe: + claude.exe: bmalloc(WebKit) → NT API(HeapAlloc) → V8 GC + DeepCode: Python dict → LRUCache → MemoryMappedIndex → IncrementalIndexer +""" + +import asyncio +import fnmatch +import hashlib +import json +import mmap +import os +import pickle +import re +import struct +import tempfile +import time +import zlib +from collections import defaultdict, deque, OrderedDict +from dataclasses import dataclass, field +from datetime import datetime +from pathlib import Path +from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union + + +# ══════════════════════════════════════════════ +# LRU 缓存 +# ══════════════════════════════════════════════ + +class LRUCache: + """ + 最近最少使用缓存 — 类似 claude.exe 中的 V8 堆缓存 + + 基于 OrderedDict 实现 O(1) 的 get/set。 + 支持: + - TTL (过期时间) + - 最大条目数限制 + - 最大内存限制 + - 访问统计 + - 逐出回调 + """ + + def __init__(self, max_size: int = 1000, max_memory_mb: int = 256, + ttl_seconds: float = 300): + self.max_size = max_size + self.max_memory_mb = max_memory_mb + self.max_memory_bytes = max_memory_mb * 1024 * 1024 + self.ttl_seconds = ttl_seconds + + self._cache: OrderedDict = OrderedDict() + self._expiry: Dict[str, float] = {} + self._size_map: Dict[str, int] = {} # key -> size in bytes + self._current_memory = 0 + self._eviction_callback: Optional[Callable] = None + + self._stats = { + "hits": 0, + "misses": 0, + "evictions": 0, + "expirations": 0, + "total_sets": 0, + "current_entries": 0, + "current_memory_mb": 0, + } + + @property + def size(self) -> int: + return len(self._cache) + + @property + def memory_usage_mb(self) -> float: + return self._current_memory / (1024 * 1024) + + def on_evict(self, callback: Callable): + """设置逐出回调""" + self._eviction_callback = callback + + def get(self, key: str, default: Any = None) -> Any: + """ + 获取缓存条目 + + Args: + key: 缓存键 + default: 默认值 + + Returns: + 缓存的值 + """ + if key not in self._cache: + self._stats["misses"] += 1 + return default + + # 检查 TTL + if key in self._expiry and time.time() > self._expiry[key]: + self._evict(key) + self._stats["expirations"] += 1 + self._stats["misses"] += 1 + return default + + # 移到末尾 (最近使用) + value = self._cache.pop(key) + self._cache[key] = value + self._stats["hits"] += 1 + return value + + def set(self, key: str, value: Any, ttl_seconds: Optional[float] = None): + """ + 设置缓存条目 + + Args: + key: 缓存键 + value: 值 + ttl_seconds: 过期时间 (覆盖默认) + """ + # 计算大小 + size = self._estimate_size(value) + + # 如果已经有该 key, 先释放内存 + if key in self._cache: + self._current_memory -= self._size_map.get(key, 0) + + # 逐出直到有足够空间 + while (self._current_memory + size > self.max_memory_bytes and self._cache): + self._evict(self._cache.keys()[0] if self._cache else None) + + # 逐出直到不超过最大条目数 + while len(self._cache) >= self.max_size: + self._evict(self._cache.keys()[0] if self._cache else None) + + self._cache[key] = value + self._size_map[key] = size + self._current_memory += size + + if ttl_seconds is not None: + self._expiry[key] = time.time() + ttl_seconds + elif self.ttl_seconds > 0: + self._expiry[key] = time.time() + self.ttl_seconds + + self._stats["total_sets"] += 1 + self._stats["current_entries"] = len(self._cache) + self._stats["current_memory_mb"] = round(self.memory_usage_mb, 2) + + def delete(self, key: str): + """删除缓存条目""" + if key in self._cache: + self._current_memory -= self._size_map.get(key, 0) + del self._cache[key] + self._expiry.pop(key, None) + self._size_map.pop(key, None) + + def clear(self): + """清空缓存""" + self._cache.clear() + self._expiry.clear() + self._size_map.clear() + self._current_memory = 0 + self._stats["current_entries"] = 0 + self._stats["current_memory_mb"] = 0 + + def contains(self, key: str) -> bool: + """检查键是否存在且未过期""" + if key not in self._cache: + return False + if key in self._expiry and time.time() > self._expiry[key]: + self._evict(key) + return False + return True + + def warmup(self, items: Dict[str, Any]): + """批量预热缓存""" + for key, value in items.items(): + self.set(key, value) + + def get_stats(self) -> dict: + total = self._stats["hits"] + self._stats["misses"] + return { + **self._stats, + "hit_rate": f"{self._stats['hits'] / max(total, 1):.1%}", + "ttl_seconds": self.ttl_seconds, + "max_size": self.max_size, + "max_memory_mb": self.max_memory_mb, + "memory_usage_pct": f"{self.memory_usage_mb / max(self.max_memory_mb, 1):.1%}", + } + + def _evict(self, key: Optional[str]): + """逐出一个条目""" + if key is None or key not in self._cache: + # 逐出最旧的 + if self._cache: + key = next(iter(self._cache)) + else: + return + + value = self._cache.pop(key, None) + size = self._size_map.pop(key, 0) + self._current_memory -= size + self._expiry.pop(key, None) + self._stats["evictions"] += 1 + + if self._eviction_callback and value is not None: + try: + self._eviction_callback(key, value) + except Exception: + pass + + def _estimate_size(self, value: Any) -> int: + """估算对象占用的字节数""" + try: + return len(pickle.dumps(value, protocol=pickle.HIGHEST_PROTOCOL)) + except Exception: + return 1024 # 默认 1KB + + +# ══════════════════════════════════════════════ +# 内存映射文件索引 +# ══════════════════════════════════════════════ + +class MemoryMappedIndex: + """ + 内存映射文件索引 — 类似 claude.exe 的 mmap 文件映射 + + 用于快速索引大型文件,无需加载完整文件到内存。 + 基于 Python 的 mmap 模块,支持: + - 文件内容搜索 (无需全量加载) + - 行级索引 (文件 → 行号 → 偏移) + - 快速 grep + - 增量更新 + """ + + def __init__(self, cache_dir: Optional[str] = None): + self.cache_dir = cache_dir or self._default_cache_dir() + os.makedirs(self.cache_dir, exist_ok=True) + + self._index: Dict[str, dict] = {} + self._mapped_files: Dict[str, mmap.mmap] = {} + self._stats = { + "files_indexed": 0, + "total_size_mb": 0, + "total_lines": 0, + "searches": 0, + "hits": 0, + "misses": 0, + "mmap_active": 0, + } + + def _default_cache_dir(self) -> str: + base = os.environ.get("DEEPCODE_CACHE_DIR", "") + if not base: + base = os.path.join(os.environ.get("HOME", os.environ.get("USERPROFILE", ".")), + ".deepcode", "cache") + return os.path.join(base, "file_index") + + def index_file(self, file_path: str) -> Optional[dict]: + """ + 索引单个文件 (构建行级索引) + + Args: + file_path: 文件路径 + + Returns: + 文件索引信息 + """ + if not os.path.exists(file_path): + return None + + abs_path = os.path.abspath(file_path) + stat = os.stat(abs_path) + + # 检查是否已被索引且未变更 + if abs_path in self._index: + existing = self._index[abs_path] + if existing.get("mtime") == stat.st_mtime and existing.get("size") == stat.st_size: + self._stats["hits"] += 1 + return existing + + # 构建索引 + try: + with open(abs_path, "r", encoding="utf-8", errors="replace") as f: + lines = [] + offset = 0 + for line_num, line in enumerate(f, 1): + lines.append({ + "line": line_num, + "offset": offset, + "length": len(line), + }) + offset += len(line) + + index = { + "path": abs_path, + "size": stat.st_size, + "mtime": stat.st_mtime, + "lines": len(lines), + "line_index": lines[:100] if len(lines) > 100 else lines, # 前100行索引 + "hash": hashlib.md5(open(abs_path, "rb").read(8192)).hexdigest(), # 头部 hash + "indexed_at": datetime.now().isoformat(), + } + + self._index[abs_path] = index + self._stats["files_indexed"] += 1 + self._stats["total_size_mb"] += stat.st_size / (1024 * 1024) + self._stats["total_lines"] += len(lines) + self._stats["misses"] += 1 + + return index + + except (IOError, UnicodeDecodeError) as e: + return None + + def index_directory(self, dir_path: str, pattern: str = "*.py") -> int: + """ + 索引目录中匹配的文件 + + Args: + dir_path: 目录路径 + pattern: glob 模式 + + Returns: + 索引的文件数 + """ + count = 0 + for root, _, files in os.walk(dir_path): + for f in files: + if fnmatch.fnmatch(f, pattern): + file_path = os.path.join(root, f) + if self.index_file(file_path): + count += 1 + return count + + def mmap_open(self, file_path: str) -> Optional[mmap.mmap]: + """ + 打开文件的内存映射 + + Args: + file_path: 文件路径 + + Returns: + mmap 对象 + """ + abs_path = os.path.abspath(file_path) + if abs_path in self._mapped_files: + return self._mapped_files[abs_path] + + if not os.path.exists(abs_path): + return None + + try: + fd = os.open(abs_path, os.O_RDONLY) + size = os.fstat(fd).st_size + if size == 0: + os.close(fd) + return None + + mapped = mmap.mmap(fd, size, access=mmap.ACCESS_READ) + os.close(fd) + self._mapped_files[abs_path] = mapped + self._stats["mmap_active"] += 1 + return mapped + except Exception: + return None + + def mmap_close(self, file_path: str): + """关闭文件的内存映射""" + abs_path = os.path.abspath(file_path) + mapped = self._mapped_files.pop(abs_path, None) + if mapped: + mapped.close() + self._stats["mmap_active"] -= 1 + + def grep(self, file_path: str, pattern: str, + max_results: int = 50) -> List[dict]: + """ + 在文件内容中搜索 (使用 mmap) + + Args: + file_path: 文件路径 + pattern: 搜索模式 (正则) + max_results: 最大结果数 + + Returns: + 匹配行列表 + """ + self._stats["searches"] += 1 + mapped = self.mmap_open(file_path) + if not mapped: + return [] + + try: + compiled = re.compile(pattern.encode("utf-8"), re.IGNORECASE) + results = [] + pos = 0 + while pos < len(mapped): + match = compiled.search(mapped, pos) + if not match: + break + # 计算行号 + line_start = mapped.rfind(b"\n", 0, match.start()) + 1 + line_end = mapped.find(b"\n", match.end()) + if line_end == -1: + line_end = len(mapped) + + line_num = mapped[:line_start].count(b"\n") + 1 + line_content = mapped[line_start:line_end].decode("utf-8", errors="replace").strip() + + results.append({ + "line": line_num, + "content": line_content[:500], + "offset": match.start(), + }) + + if len(results) >= max_results: + break + pos = line_end + 1 + + return results + + except (re.error, Exception) as e: + return [] + finally: + self.mmap_close(file_path) + + def close_all(self): + """关闭所有内存映射""" + for path in list(self._mapped_files.keys()): + self.mmap_close(path) + + def get_stats(self) -> dict: + return dict(self._stats) + + def __del__(self): + self.close_all() + + +# ══════════════════════════════════════════════ +# 增量索引器 +# ══════════════════════════════════════════════ + +class IncrementalIndexer: + """ + 增量索引器 — 只索引变更的文件 + + 对比 claude.exe 的全量索引,DeepCode 只需索引变更文件。 + 使用文件时间戳 + 内容 hash 检测变更。 + + 用法: + indexer = IncrementalIndexer() + indexer.watch_directory("./src") # 开始监控 + changed = indexer.get_changed_files() # 获取变更 + indexer.update_index(changed) # 增量更新 + """ + + def __init__(self, cache_dir: Optional[str] = None): + self.cache_dir = cache_dir or os.path.join( + os.environ.get("DEEPCODE_CACHE_DIR", ""), + ".deepcode", "cache", "incremental_index" + ) + os.makedirs(self.cache_dir, exist_ok=True) + + self._file_hashes: Dict[str, str] = self._load_hashes() + self._watched_dirs: Set[str] = set() + self._index_data: Dict[str, dict] = {} + self._stats = { + "total_files": 0, + "indexed_files": 0, + "changed_files": 0, + "new_files": 0, + "deleted_files": 0, + "full_rebuilds": 0, + } + + def watch_directory(self, dir_path: str): + """添加监控目录""" + abs_path = os.path.abspath(dir_path) + if os.path.isdir(abs_path): + self._watched_dirs.add(abs_path) + + def unwatch_directory(self, dir_path: str): + """移除监控目录""" + self._watched_dirs.discard(os.path.abspath(dir_path)) + + def get_changed_files(self, pattern: str = "*") -> Dict[str, List[str]]: + """ + 检测变更的文件 + + Returns: + {"new": [...], "changed": [...], "deleted": [...]} + """ + result = {"new": [], "changed": [], "deleted": []} + + # 收集当前文件 + current_files = set() + for watched_dir in self._watched_dirs: + if not os.path.isdir(watched_dir): + continue + for root, _, files in os.walk(watched_dir): + for f in files: + if fnmatch.fnmatch(f, pattern): + current_files.add(os.path.join(root, f)) + + # 检测新增和变更 + for file_path in current_files: + current_hash = self._hash_file(file_path) + if current_hash is None: + continue + + prev_hash = self._file_hashes.get(file_path) + if prev_hash is None: + result["new"].append(file_path) + elif prev_hash != current_hash: + result["changed"].append(file_path) + + # 检测删除 + for file_path in list(self._file_hashes.keys()): + if file_path not in current_files: + result["deleted"].append(file_path) + + self._stats["changed_files"] = len(result["changed"]) + self._stats["new_files"] = len(result["new"]) + self._stats["deleted_files"] = len(result["deleted"]) + + return result + + def update_index(self, changes: Dict[str, List[str]], + index_func: Callable[[str], Any] = None): + """ + 增量更新索引 + + Args: + changes: get_changed_files 的返回值 + index_func: 索引函数 (接收文件路径, 返回索引数据) + """ + # 处理新增和变更 + for file_path in changes.get("new", []) + changes.get("changed", []): + if index_func: + try: + self._index_data[file_path] = index_func(file_path) + except Exception: + pass + current_hash = self._hash_file(file_path) + if current_hash: + self._file_hashes[file_path] = current_hash + self._stats["indexed_files"] += 1 + + # 处理删除 + for file_path in changes.get("deleted", []): + self._file_hashes.pop(file_path, None) + self._index_data.pop(file_path, None) + + self._stats["total_files"] = len(self._file_hashes) + self._save_hashes() + + def needs_full_rebuild(self) -> bool: + """判断是否需要全量重建""" + if self._stats["changed_files"] > self._stats["total_files"] * 0.5: + return True + if not self._file_hashes: + return True + return False + + def full_rebuild(self, index_func: Callable[[str], Any], + dir_path: str, pattern: str = "*"): + """全量重建索引""" + self._file_hashes.clear() + self._index_data.clear() + + for root, _, files in os.walk(dir_path): + for f in files: + if fnmatch.fnmatch(f, pattern): + file_path = os.path.join(root, f) + if index_func: + try: + self._index_data[file_path] = index_func(file_path) + except Exception: + pass + current_hash = self._hash_file(file_path) + if current_hash: + self._file_hashes[file_path] = current_hash + + self._stats["full_rebuilds"] += 1 + self._stats["total_files"] = len(self._file_hashes) + self._stats["indexed_files"] = len(self._index_data) + self._save_hashes() + + def query(self, file_path: str) -> Optional[dict]: + """查询文件的索引数据""" + return self._index_data.get(os.path.abspath(file_path)) + + def _hash_file(self, file_path: str) -> Optional[str]: + """计算文件内容 hash""" + try: + with open(file_path, "rb") as f: + # 只读前 4KB 和最后 1KB 用于快速变更检测 + head = f.read(4096) + f.seek(-1024, os.SEEK_END) + tail = f.read(1024) + return hashlib.md5(head + tail).hexdigest() + except (IOError, OSError): + return None + + def _save_hashes(self): + """持久化 hash 表""" + hash_path = os.path.join(self.cache_dir, "file_hashes.json") + os.makedirs(os.path.dirname(hash_path), exist_ok=True) + try: + with open(hash_path, "w", encoding="utf-8") as f: + json.dump(self._file_hashes, f, indent=2) + except Exception: + pass + + def _load_hashes(self) -> Dict[str, str]: + """加载 hash 表""" + hash_path = os.path.join(self.cache_dir, "file_hashes.json") + if os.path.exists(hash_path): + try: + with open(hash_path, "r", encoding="utf-8") as f: + return json.load(f) + except (json.JSONDecodeError, IOError): + pass + return {} + + def get_stats(self) -> dict: + return dict(self._stats) + + +# ══════════════════════════════════════════════ +# Token 预算管理器 +# ══════════════════════════════════════════════ + +class TokenBudgetManager: + """ + Token 预算管理器 — 控制每轮/每小时/每天的 Token 消耗 + + 类似 claude.exe 的 token_tracker.ts + budget_controller.ts + + 用法: + budget = TokenBudgetManager() + budget.use(1500) # 消耗 1500 tokens + if budget.remaining < 1000: + print("预算不足") + """ + + def __init__(self, max_per_turn: int = 8000, + max_per_hour: int = 100000, + max_per_day: int = 1000000): + self.max_per_turn = max_per_turn + self.max_per_hour = max_per_hour + self.max_per_day = max_per_day + + self._turn_used = 0 + self._hour_used = 0 + self._day_used = 0 + self._hour_reset = time.time() + self._day_reset = time.time() + self._history: deque = deque(maxlen=1000) + self._total_used = 0 + self._warning_threshold = 0.85 + self._stats = {"warnings": 0, "blocks": 0} + + def start_turn(self): + """开始新一轮 (重置轮次计数)""" + self._turn_used = 0 + + def use(self, tokens: int) -> bool: + """ + 消耗 tokens + + Args: + tokens: 消耗数量 + + Returns: + True: 消耗成功 + False: 超过预算 + """ + now = time.time() + + # 重置 + if now - self._hour_reset >= 3600: + self._hour_used = 0 + self._hour_reset = now + if now - self._day_reset >= 86400: + self._day_used = 0 + self._day_reset = now + + # 检查预算 + if self._turn_used + tokens > self.max_per_turn: + self._stats["blocks"] += 1 + return False + if self._hour_used + tokens > self.max_per_hour: + self._stats["blocks"] += 1 + return False + if self._day_used + tokens > self.max_per_day: + self._stats["blocks"] += 1 + return False + + self._turn_used += tokens + self._hour_used += tokens + self._day_used += tokens + self._total_used += tokens + + # 警告 + if self.usage_pct >= self._warning_threshold: + self._stats["warnings"] += 1 + + self._history.append({ + "tokens": tokens, + "timestamp": datetime.now().isoformat(), + }) + + return True + + @property + def remaining(self) -> int: + return max(0, self.max_per_turn - self._turn_used) + + @property + def usage_pct(self) -> float: + return self._turn_used / max(self.max_per_turn, 1) + + def summary(self) -> dict: + return { + "per_turn": {"max": self.max_per_turn, "used": self._turn_used, + "remaining": self.remaining, + "pct": f"{self.usage_pct:.1%}"}, + "per_hour": {"max": self.max_per_hour, "used": self._hour_used}, + "per_day": {"max": self.max_per_day, "used": self._day_used}, + "total_used": self._total_used, + "warning_threshold": f"{self._warning_threshold:.0%}", + **self._stats, + } + + +# ══════════════════════════════════════════════ +# 缓存预热器 +# ══════════════════════════════════════════════ + +class CacheWarmup: + """ + 缓存预热器 — 启动时预加载常用数据 + + 参考 claude.exe 的 V8 快照热身机制: + - 启动时预编译常用模块 + - 预加载常用文件到缓存 + """ + + def __init__(self, cache: LRUCache): + self.cache = cache + self._stats = {"warmed_items": 0, "errors": 0} + + async def warmup_files(self, file_paths: List[str]): + """预热文件缓存""" + for path in file_paths: + try: + if os.path.exists(path) and os.path.isfile(path): + with open(path, "rb") as f: + content = f.read(8192) # 只缓存前 8KB + self.cache.set(f"file_head:{path}", content, ttl_seconds=600) + self._stats["warmed_items"] += 1 + except Exception: + self._stats["errors"] += 1 + + async def warmup_common_queries(self): + """预热常用查询""" + common = { + "tools:list": {"tools": []}, # 占位 + "system:status": {"status": "ready"}, + } + for key, value in common.items(): + self.cache.set(key, value, ttl_seconds=3600) + self._stats["warmed_items"] += 1 + + def get_stats(self) -> dict: + return dict(self._stats) + + +# ══════════════════════════════════════════════ +# NetnodeStore — IDA Pro 风格持久化 KV 存储 +# ══════════════════════════════════════════════ + +class NetnodeStore: + """ + IDA Pro 风格的 netnode 持久化 KV 存储 — 参考 ida.dll 的 68 个 netnode_* API + + IDA 的 netnode 系统是其数据库 (IDB/I64) 的核心。每个 netnode 是一个 + 持久化的键值节点,支持: + - 唯一名称/ID 标识 + - 主值 (value) + 数值辅助值 (altval) + - 任意数量的辅助键值对 (supval) + - blob 二进制数据 + - 哈希索引 (hashval) + - 父子关系 (通过 supval 实现) + + DEEPCODE 实现使用 SQLite 作为持久化后端, + 在内存中缓存热节点 (通过 LRUCache)。 + + 用法: + store = NetnodeStore("analysis_results") + node = store.create("func_1000_analysis") + node.set_value({"name": "sub_1000", "size": 128}) + node.set_supval("callers", ["main", "start"]) + node.set_blob(b"raw data...") + store.flush() + """ + + def __init__(self, namespace: str = "default", + db_path: Optional[str] = None, + lru_size: int = 500): + self.namespace = namespace + self._lru = LRUCache(max_size=lru_size, max_memory_mb=64) + self._db_path = db_path or self._default_db_path() + self._dirty_nodes: Set[str] = set() + self._stats = { + "nodes_created": 0, + "nodes_loaded": 0, + "values_set": 0, + "values_get": 0, + "blobs_set": 0, + "blobs_get": 0, + "supvals_set": 0, + "supvals_get": 0, + "flushes": 0, + "hash_hits": 0, + "hash_misses": 0, + } + self._conn = None + self._init_db() + + def _default_db_path(self) -> str: + base = os.environ.get("DEEPCODE_DATA_DIR", "") + if not base: + base = os.path.join( + os.environ.get("HOME", os.environ.get("USERPROFILE", ".")), + ".deepcode", "data" + ) + os.makedirs(base, exist_ok=True) + return os.path.join(base, f"netnode_{self.namespace}.db") + + def _init_db(self): + import sqlite3 + db_dir = os.path.dirname(self._db_path) + if db_dir: + os.makedirs(db_dir, exist_ok=True) + self._conn = sqlite3.connect(self._db_path, timeout=10) + conn = self._conn + conn.execute(""" + CREATE TABLE IF NOT EXISTS netnodes ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT UNIQUE NOT NULL, + value BLOB, + altval INTEGER DEFAULT 0, + created_at TEXT DEFAULT (datetime('now')), + updated_at TEXT DEFAULT (datetime('now')) + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS supvals ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + node_id INTEGER NOT NULL, + key TEXT NOT NULL, + value BLOB, + value_type TEXT DEFAULT 'text', + FOREIGN KEY (node_id) REFERENCES netnodes(id) ON DELETE CASCADE, + UNIQUE(node_id, key) + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS blobs ( + node_id INTEGER PRIMARY KEY, + data BLOB, + size INTEGER DEFAULT 0, + md5 TEXT, + FOREIGN KEY (node_id) REFERENCES netnodes(id) ON DELETE CASCADE + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS node_relations ( + parent_id INTEGER NOT NULL, + child_id INTEGER NOT NULL, + rel_type TEXT DEFAULT 'default', + PRIMARY KEY (parent_id, child_id), + FOREIGN KEY (parent_id) REFERENCES netnodes(id) ON DELETE CASCADE, + FOREIGN KEY (child_id) REFERENCES netnodes(id) ON DELETE CASCADE + ) + """) + conn.execute(""" + CREATE INDEX IF NOT EXISTS idx_netnodes_name ON netnodes(name) + """) + conn.commit() + + def _get_conn(self): + return self._conn + + def create(self, name: str, value: Any = None, altval: int = 0) -> Optional[int]: + import sqlite3 + conn = self._get_conn() + try: + val_bytes = self._serialize(value) if value is not None else None + conn.execute( + "INSERT OR IGNORE INTO netnodes (name, value, altval) VALUES (?, ?, ?)", + (name, val_bytes, altval) + ) + conn.commit() + cursor = conn.execute("SELECT id FROM netnodes WHERE name = ?", (name,)) + row = cursor.fetchone() + if row: + self._stats["nodes_created"] += 1 + self._invalidate_cache(name) + return row[0] + return None + except sqlite3.IntegrityError: + return None + finally: + pass + + def get_node_id(self, name: str) -> Optional[int]: + cache_key = f"node_id:{name}" + cached = self._lru.get(cache_key) + if cached is not None: + self._stats["hash_hits"] += 1 + return cached + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute("SELECT id FROM netnodes WHERE name = ?", (name,)) + row = cursor.fetchone() + if row: + self._stats["hash_misses"] += 1 + self._lru.set(cache_key, row[0], ttl_seconds=3600) + return row[0] + return None + finally: + pass + + def delete(self, name: str) -> bool: + import sqlite3 + node_id = self.get_node_id(name) + if node_id is None: + return False + conn = self._get_conn() + try: + conn.execute("DELETE FROM blobs WHERE node_id = ?", (node_id,)) + conn.execute("DELETE FROM supvals WHERE node_id = ?", (node_id,)) + conn.execute( + "DELETE FROM node_relations WHERE parent_id = ? OR child_id = ?", + (node_id, node_id) + ) + conn.execute("DELETE FROM netnodes WHERE id = ?", (node_id,)) + conn.commit() + self._invalidate_cache(name) + return True + finally: + pass + + def exists(self, name: str) -> bool: + return self.get_node_id(name) is not None + + def set_value(self, name: str, value: Any) -> bool: + import sqlite3 + node_id = self.get_node_id(name) + if node_id is None: + node_id = self.create(name) + if node_id is None: + return False + conn = self._get_conn() + try: + val_bytes = self._serialize(value) + conn.execute("UPDATE netnodes SET value = ?, updated_at = datetime('now') WHERE id = ?", + (val_bytes, node_id)) + conn.commit() + self._stats["values_set"] += 1 + self._invalidate_cache(name) + return True + finally: + pass + + def get_value(self, name: str, default: Any = None) -> Any: + cache_key = f"value:{name}" + cached = self._lru.get(cache_key) + if cached is not None: + return cached + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute("SELECT value FROM netnodes WHERE name = ?", (name,)) + row = cursor.fetchone() + if row and row[0]: + result = self._deserialize(row[0]) + self._lru.set(cache_key, result, ttl_seconds=300) + return result + return default + finally: + pass + + def set_altval(self, name: str, altval: int) -> bool: + import sqlite3 + node_id = self.get_node_id(name) + if node_id is None: + node_id = self.create(name) + if node_id is None: + return False + conn = self._get_conn() + try: + conn.execute("UPDATE netnodes SET altval = ?, updated_at = datetime('now') WHERE id = ?", + (altval, node_id)) + conn.commit() + return True + finally: + pass + + def get_altval(self, name: str) -> int: + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute("SELECT altval FROM netnodes WHERE name = ?", (name,)) + row = cursor.fetchone() + return row[0] if row else 0 + finally: + pass + + def set_supval(self, name: str, key: str, value: Any, value_type: str = "auto") -> bool: + import sqlite3 + node_id = self.get_node_id(name) + if node_id is None: + node_id = self.create(name) + if node_id is None: + return False + if value_type == "auto": + value_type = "text" if isinstance(value, (int, float, bool, str)) else "json" + conn = self._get_conn() + try: + conn.execute(""" + INSERT OR REPLACE INTO supvals (node_id, key, value, value_type) + VALUES (?, ?, ?, ?) + """, (node_id, key, self._serialize(value), value_type)) + conn.commit() + self._stats["supvals_set"] += 1 + self._invalidate_cache(f"supval:{name}:{key}") + return True + finally: + pass + + def get_supval(self, name: str, key: str, default: Any = None) -> Any: + cache_key = f"supval:{name}:{key}" + cached = self._lru.get(cache_key) + if cached is not None: + return cached + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute(""" + SELECT s.value FROM supvals s + JOIN netnodes n ON s.node_id = n.id + WHERE n.name = ? AND s.key = ? + """, (name, key)) + row = cursor.fetchone() + if row: + result = self._deserialize(row[0]) + self._lru.set(cache_key, result, ttl_seconds=300) + return result + return default + finally: + pass + + def get_all_supvals(self, name: str) -> Dict[str, Any]: + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute(""" + SELECT s.key, s.value FROM supvals s + JOIN netnodes n ON s.node_id = n.id + WHERE n.name = ? + """, (name,)) + return {row[0]: self._deserialize(row[1]) for row in cursor.fetchall()} + finally: + pass + + def del_supval(self, name: str, key: str) -> bool: + import sqlite3 + node_id = self.get_node_id(name) + if node_id is None: + return False + conn = self._get_conn() + try: + conn.execute("DELETE FROM supvals WHERE node_id = ? AND key = ?", (node_id, key)) + conn.commit() + self._invalidate_cache(f"supval:{name}:{key}") + return True + finally: + pass + + def set_blob(self, name: str, data: bytes) -> bool: + import sqlite3 + node_id = self.get_node_id(name) + if node_id is None: + node_id = self.create(name) + if node_id is None: + return False + md5_hash = hashlib.md5(data).hexdigest() + conn = self._get_conn() + try: + conn.execute(""" + INSERT OR REPLACE INTO blobs (node_id, data, size, md5) + VALUES (?, ?, ?, ?) + """, (node_id, data, len(data), md5_hash)) + conn.commit() + self._stats["blobs_set"] += 1 + self._invalidate_cache(f"blob:{name}") + return True + finally: + pass + + def get_blob(self, name: str) -> Optional[bytes]: + cache_key = f"blob:{name}" + cached = self._lru.get(cache_key) + if cached is not None: + return cached + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute(""" + SELECT b.data FROM blobs b + JOIN netnodes n ON b.node_id = n.id WHERE n.name = ? + """, (name,)) + row = cursor.fetchone() + if row: + self._lru.set(cache_key, row[0], ttl_seconds=600) + return row[0] + return None + finally: + pass + + def get_blob_info(self, name: str) -> Optional[dict]: + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute(""" + SELECT b.size, b.md5, n.updated_at FROM blobs b + JOIN netnodes n ON b.node_id = n.id WHERE n.name = ? + """, (name,)) + row = cursor.fetchone() + if row: + return {"size": row[0], "md5": row[1], "updated_at": row[2]} + return None + finally: + pass + + def add_child(self, parent_name: str, child_name: str, rel_type: str = "default") -> bool: + import sqlite3 + parent_id = self.get_node_id(parent_name) + child_id = self.get_node_id(child_name) + if parent_id is None: + parent_id = self.create(parent_name) + if child_id is None: + child_id = self.create(child_name) + if parent_id is None or child_id is None: + return False + conn = self._get_conn() + try: + conn.execute( + "INSERT OR IGNORE INTO node_relations (parent_id, child_id, rel_type) VALUES (?, ?, ?)", + (parent_id, child_id, rel_type) + ) + conn.commit() + return True + finally: + pass + + def get_children(self, name: str, rel_type: Optional[str] = None) -> List[str]: + import sqlite3 + conn = self._get_conn() + try: + if rel_type: + cursor = conn.execute(""" + SELECT n.name FROM netnodes n + JOIN node_relations r ON n.id = r.child_id + JOIN netnodes p ON p.id = r.parent_id + WHERE p.name = ? AND r.rel_type = ? + """, (name, rel_type)) + else: + cursor = conn.execute(""" + SELECT n.name FROM netnodes n + JOIN node_relations r ON n.id = r.child_id + JOIN netnodes p ON p.id = r.parent_id + WHERE p.name = ? + """, (name,)) + return [row[0] for row in cursor.fetchall()] + finally: + pass + + def get_parents(self, name: str) -> List[str]: + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute(""" + SELECT p.name FROM netnodes p + JOIN node_relations r ON p.id = r.parent_id + JOIN netnodes c ON c.id = r.child_id WHERE c.name = ? + """, (name,)) + return [row[0] for row in cursor.fetchall()] + finally: + pass + + def search_by_supval(self, key: str, value: Any) -> List[str]: + import sqlite3 + conn = self._get_conn() + try: + cursor = conn.execute(""" + SELECT n.name FROM netnodes n + JOIN supvals s ON n.id = s.node_id + WHERE s.key = ? AND s.value = ? + """, (key, self._serialize(value))) + return [row[0] for row in cursor.fetchall()] + finally: + pass + + def search_by_name_glob(self, pattern: str) -> List[str]: + import sqlite3 + conn = self._get_conn() + try: + sql_pattern = pattern.replace("*", "%").replace("?", "_") + cursor = conn.execute( + "SELECT name FROM netnodes WHERE name LIKE ? ORDER BY name", + (sql_pattern,) + ) + return [row[0] for row in cursor.fetchall()] + finally: + pass + + def flush(self): + self._dirty_nodes.clear() + self._stats["flushes"] += 1 + + def get_size(self) -> dict: + import sqlite3 + conn = self._get_conn() + try: + nc = conn.execute("SELECT COUNT(*) FROM netnodes").fetchone()[0] + sc = conn.execute("SELECT COUNT(*) FROM supvals").fetchone()[0] + bc = conn.execute("SELECT COUNT(*) FROM blobs").fetchone()[0] + bs = conn.execute("SELECT COALESCE(SUM(size), 0) FROM blobs").fetchone()[0] + return {"nodes": nc, "supvals": sc, "blobs": bc, "blob_size_bytes": bs} + finally: + pass + + def get_stats(self) -> dict: + return dict(self._stats) + + def vacuum(self): + import sqlite3 + conn = self._get_conn() + try: + conn.execute("VACUUM") + finally: + pass + + def _serialize(self, value: Any) -> bytes: + try: + return pickle.dumps(value, protocol=pickle.HIGHEST_PROTOCOL) + except Exception: + return str(value).encode("utf-8") + + def _deserialize(self, data: bytes) -> Any: + if data is None: + return None + try: + return pickle.loads(data) + except Exception: + try: + return data.decode("utf-8") + except Exception: + return data + + def _invalidate_cache(self, name: str): + self._lru.delete(f"node_id:{name}") + self._lru.delete(f"value:{name}") + self._lru.delete(f"blob:{name}") + self._dirty_nodes.add(name) diff --git a/mcp_servers/signature_engine.py b/mcp_servers/signature_engine.py new file mode 100644 index 000000000..a4bd7a11c --- /dev/null +++ b/mcp_servers/signature_engine.py @@ -0,0 +1,626 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +SignatureEngine v1.0 — IDA FLIRT 风格函数签名引擎 +================================================== +参考: IDA Pro 9.2 的 FLIRT (Fast Library Identification and Recognition Technology) + +FLIRT 的核心设计: + 1. 函数起始字节的 CRC16 签名 (胖签名: 前 32-64 字节) + 2. 签名数据库按 CRC16 排序,支持二分查找 + 3. 每个签名关联: 函数名、库名、调用约定、参数信息 + 4. 匹配后自动重命名/注释函数 + +DEEPCODE 实现: + - CRC16/CRC32 双模式签名生成 + - SQLite 签名数据库 (通过 NetnodeStore) + - 预置 Rust 标准库签名 (从 deepcode-rust-re 导入) + - Ghidra MCP 集成 (批量扫描函数) + - 自动重命名/注释匹配的函数 + +用法: + engine = SignatureEngine() + engine.load_library("rust_std") + matches = engine.scan_functions(functions_batch) + engine.apply_matches(matches) +""" + +import hashlib +import json +import os +import struct +import time +import zlib +from typing import Any, Callable, Dict, List, Optional, Set, Tuple + +# ── CRC16 表 (CRC-16/ARC) ── + +CRC16_TABLE = None + +def _build_crc16_table() -> List[int]: + table = [] + for i in range(256): + crc = i + for _ in range(8): + if crc & 1: + crc = (crc >> 1) ^ 0xA001 + else: + crc >>= 1 + table.append(crc) + return table + +def crc16(data: bytes) -> int: + """计算 CRC16 校验和 (CRC-16/ARC)""" + global CRC16_TABLE + if CRC16_TABLE is None: + CRC16_TABLE = _build_crc16_table() + crc = 0 + for byte in data: + crc = (crc >> 8) ^ CRC16_TABLE[(crc ^ byte) & 0xFF] + return crc + +def crc32(data: bytes) -> int: + """计算 CRC32 校验和""" + return zlib.crc32(data) & 0xFFFFFFFF + + +# ══════════════════════════════════════════════ +# 签名条目 +# ══════════════════════════════════════════════ + +class Signature: + """ + 单个函数签名 — 对应 IDA FLIRT 的一条签名记录 + + 属性: + crc: CRC16 (前 32 字节) — 快速预过滤 + crc_full: CRC32 (前 64 字节) — 精确匹配 + name: 函数名 + library: 所属库 + version: 库版本 + pattern: 原始字节模板 (带掩码) + mask: xx?x 格式的掩码 (可选字节) + size: 签名覆盖的字节数 + info: 额外元信息 + """ + + __slots__ = ('crc', 'crc_full', 'name', 'library', 'version', + 'pattern', 'mask', 'size', 'info') + + def __init__(self, name: str, library: str = "unknown", + version: str = "", pattern: Optional[bytes] = None, + mask: Optional[str] = None): + self.crc = 0 + self.crc_full = 0 + self.name = name + self.library = library + self.version = version + self.pattern = pattern + self.mask = mask + self.size = len(pattern) if pattern else 0 + self.info: Dict[str, Any] = {} + + if pattern: + head32 = pattern[:32] + head64 = pattern[:64] + self.crc = crc16(head32) if head32 else 0 + self.crc_full = crc32(head64) if head64 else 0 + + def to_dict(self) -> dict: + return { + "crc": self.crc, + "crc_full": self.crc_full, + "name": self.name, + "library": self.library, + "version": self.version, + "size": self.size, + "pattern": self.pattern.hex() if self.pattern else None, + "mask": self.mask, + "info": self.info, + } + + @classmethod + def from_dict(cls, d: dict) -> "Signature": + sig = cls( + name=d["name"], + library=d.get("library", "unknown"), + version=d.get("version", ""), + pattern=bytes.fromhex(d["pattern"]) if d.get("pattern") else None, + mask=d.get("mask"), + ) + sig.crc = d.get("crc", 0) + sig.crc_full = d.get("crc_full", 0) + sig.size = d.get("size", 0) + sig.info = d.get("info", {}) + return sig + + def __repr__(self) -> str: + return f"" + + +# ══════════════════════════════════════════════ +# 签名库 +# ══════════════════════════════════════════════ + +class SignatureLibrary: + """ + 签名库 — 一组按 CRC16 排序的签名集合 (对应 IDA FLIRT .sig 文件) + + 支持: + - 二分查找快速匹配 (O(log n)) + - 批量导入/导出 + - 合并/增量更新 + """ + + def __init__(self, name: str = "default"): + self.name = name + self._signatures: Dict[int, List[Signature]] = {} # crc -> [sig] + self._crc_full_index: Dict[int, Signature] = {} # crc_full -> sig + self._stats = { + "total": 0, + "libraries": set(), + "match_attempts": 0, + "match_hits": 0, + } + + @property + def total(self) -> int: + return self._stats["total"] + + def add(self, sig: Signature): + """添加一条签名到库""" + if sig.crc not in self._signatures: + self._signatures[sig.crc] = [] + self._signatures[sig.crc].append(sig) + if sig.crc_full: + self._crc_full_index[sig.crc_full] = sig + self._stats["total"] += 1 + self._stats["libraries"].add(sig.library) + + def add_batch(self, sigs: List[Signature]): + """批量添加签名""" + for sig in sigs: + self.add(sig) + + def match(self, head32: bytes, head64: Optional[bytes] = None, + min_confidence: float = 0.6) -> Optional[Signature]: + """ + 匹配函数签名 + + Args: + head32: 函数前 32 字节 + head64: 函数前 64 字节 (可选,提高精度) + min_confidence: 最低置信度 + + Returns: + 匹配的签名或 None + """ + self._stats["match_attempts"] += 1 + if not head32: + return None + + crc_val = crc16(head32) + candidates = self._signatures.get(crc_val, []) + + if not candidates: + return None + + # 如果有 CRC32 匹配,精确命中 + if head64 and len(candidates) > 1: + crc_full_val = crc32(head64) + exact = self._crc_full_index.get(crc_full_val) + if exact: + self._stats["match_hits"] += 1 + return exact + + # 单候选,直接返回 + if len(candidates) == 1: + self._stats["match_hits"] += 1 + return candidates[0] + + # 多候选: 优先按掩码匹配 + best = None + best_score = 0 + for sig in candidates: + if sig.pattern and sig.mask: + score = self._mask_match(head32, sig.pattern, sig.mask) + if score > best_score: + best_score = score + best = sig + + if best and best_score >= min_confidence: + self._stats["match_hits"] += 1 + return best + + # 无掩码匹配时: 返回长度最匹配的候选 (best-effort) + if candidates: + self._stats["match_hits"] += 1 + return candidates[0] + + return None + + def match_batch(self, functions: List[dict]) -> List[dict]: + """ + 批量匹配函数 + + Args: + functions: [{"name": str, "head32": bytes, "head64": bytes}, ...] + + Returns: + [{"func_name": str, "signature": Signature, "confidence": float}, ...] + """ + results = [] + for func in functions: + sig = self.match( + func.get("head32", b""), + func.get("head64"), + ) + if sig: + results.append({ + "func_name": func.get("name", ""), + "address": func.get("address", ""), + "signature": sig, + "library": sig.library, + "matched_name": sig.name, + }) + return results + + def _mask_match(self, data: bytes, pattern: bytes, + mask: str) -> float: + """带掩码的字节匹配,返回 0.0-1.0 的匹配度""" + if len(data) != len(pattern): + return 0.0 + matches = 0 + total = 0 + for i, (d, p, m) in enumerate(zip(data, pattern, mask)): + if m == 'x': # 必须匹配 + total += 1 + if d == p: + matches += 1 + # '?' = 忽略 + return matches / max(total, 1) + + def get_stats(self) -> dict: + return { + "name": self.name, + "total_signatures": self._stats["total"], + "libraries": sorted(self._stats["libraries"]), + "unique_crcs": len(self._signatures), + "match_attempts": self._stats["match_attempts"], + "match_hits": self._stats["match_hits"], + "hit_rate": f"{self._stats['match_hits'] / max(self._stats['match_attempts'], 1):.1%}", + } + + def export_json(self) -> str: + """导出为 JSON 格式""" + sigs = [] + for crc, sig_list in self._signatures.items(): + for sig in sig_list: + sigs.append(sig.to_dict()) + return json.dumps({ + "library": self.name, + "total": len(sigs), + "signatures": sigs, + }, indent=2, ensure_ascii=False) + + @classmethod + def import_json(cls, json_str: str) -> "SignatureLibrary": + """从 JSON 导入""" + data = json.loads(json_str) + lib = cls(name=data.get("library", "imported")) + for sig_data in data.get("signatures", []): + lib.add(Signature.from_dict(sig_data)) + return lib + + +# ══════════════════════════════════════════════ +# 签名引擎 +# ══════════════════════════════════════════════ + +class SignatureEngine: + """ + FLIRT 风格签名引擎 — 综合签名匹配系统 + + 整合: + - 签名库管理 (SignatureLibrary) + - 签名生成 (从函数字节/汇编/反编译模型) + - 批量扫描 (通过 Ghidra MCP) + - NetnodeStore 持久化 + - Rust RE 技能集成 + + 用法: + engine = SignatureEngine() + engine.add_builtin_rules() # 加载内置库规则 + engine.load_from_store("my_sigs") # 从持久化存储加载 + results = engine.scan([func1, func2]) + engine.apply_to_ghidra(results) # 可选: 推送到 Ghidra + """ + + def __init__(self, store: Optional[Any] = None): + self._libraries: Dict[str, SignatureLibrary] = {} + self._store = store # NetnodeStore 实例 + self._crc_cache: Dict[int, str] = {} # crc -> library_name + self._stats = { + "total_libraries": 0, + "total_signatures": 0, + "total_scans": 0, + "total_matches": 0, + } + + # ── 库管理 ── + + def add_library(self, lib: SignatureLibrary): + """添加签名库""" + self._libraries[lib.name] = lib + self._stats["total_libraries"] = len(self._libraries) + self._stats["total_signatures"] = sum( + lib._stats["total"] for lib in self._libraries.values() + ) + # 更新 CRC 缓存 + for crc in lib._signatures: + self._crc_cache[crc] = lib.name + + def get_library(self, name: str) -> Optional[SignatureLibrary]: + return self._libraries.get(name) + + def list_libraries(self) -> List[str]: + return sorted(self._libraries.keys()) + + # ── 内置规则 ── + + def add_builtin_rules(self): + """ + 加载内置库签名规则 + 这些是已知的编译时特征和常见库的识别模式 + """ + lib = SignatureLibrary("builtin_patterns") + + # Rust 标准库识别模式 (基于 IDA FLIRT Rust 签名) + rust_sigs = self._load_rust_signatures() + lib.add_batch(rust_sigs) + + # MSVC CRT 签名 + msvc_sigs = self._load_msvc_signatures() + lib.add_batch(msvc_sigs) + + self.add_library(lib) + return lib + + def _load_rust_signatures(self) -> List[Signature]: + """从 Rust RE skill 的知识生成 Rust 标准库签名""" + sigs = [] + + # Rust panic 相关函数 (特征字节序列) + rust_patterns = { + "core::panicking::panic": b"\x48\x83\xec\x28\xe8", + "core::panicking::panic_fmt": b"\x48\x89\x5c\x24\x08\x57\x48\x83\xec\x30", + "core::result::unwrap_failed": b"\x48\x83\xec\x28\x48\x8b\x05", + "std::sys::windows::alloc::System::alloc": b"\x48\x89\x5c\x24\x08\x57\x48\x83\xec\x20", + "std::rt::lang_start": b"\x48\x83\xec\x38\x48\x8b\x05", + "std::rt::lang_start_internal": b"\x55\x57\x41\x54\x41\x55\x41\x56", + "core::ptr::drop_in_place": b"\x48\x83\xec\x28\xe8", + "alloc::alloc::alloc": b"\x48\x89\x5c\x24\x08\x57\x48\x83\xec\x20\x48\x8b\xd9", + } + for name, pattern in rust_patterns.items(): + sig = Signature(name, library="rust_std", pattern=pattern) + sig.info["source"] = "builtin_rust" + sigs.append(sig) + + # Tokio 运行时 + tokio_patterns = { + "tokio::runtime::Runtime::block_on": b"\x48\x89\x5c\x24\x08\x57\x48\x83\xec\x20\x48\x8b\xd9\xe8", + "tokio::runtime::Runtime::new": b"\x48\x83\xec\x28\xe8", + "tokio::spawn": b"\x48\x89\x5c\x24\x08\x57\x48\x83\xec\x20", + } + for name, pattern in tokio_patterns.items(): + sig = Signature(name, library="tokio", pattern=pattern) + sig.info["source"] = "builtin_tokio" + sigs.append(sig) + + return sigs + + def _load_msvc_signatures(self) -> List[Signature]: + """MSVC CRT 标准函数签名""" + sigs = [] + msvc_patterns = { + "memset": b"\x48\x83\xec\x28\x48\x85\xd2\x74", + "memcpy": b"\x48\x83\xec\x28\x48\x85\xd2\x74", + "memmove": b"\x48\x83\xec\x28\x4c\x8b\xdc", + "strlen": b"\x48\x83\xec\x28\x48\x85\xc9\x74", + "malloc": b"\x48\x83\xec\x28\x65\x48\x8b\x04\x25", + "free": b"\x48\x83\xec\x28\x48\x85\xc9\x74", + "calloc": b"\x48\x83\xec\x28\x45\x33\xc0", + "realloc": b"\x48\x83\xec\x28\x48\x85\xd2\x74", + "printf": b"\x48\x83\xec\x28\x48\x8b\xc2", + "sprintf": b"\x48\x89\x5c\x24\x08\x48\x89\x74\x24\x10", + "fopen": b"\x48\x83\xec\x28\x48\x85\xd2\x74", + "fread": b"\x48\x83\xec\x28\x45\x85\xc0", + "fwrite": b"\x48\x83\xec\x28\x4d\x85\xc0", + "qsort": b"\x48\x83\xec\x28\x4c\x8b\xdc", + "bsearch": b"\x48\x83\xec\x28\x48\x85\xd2\x74", + } + for name, pattern in msvc_patterns.items(): + sig = Signature(name, library="msvcrt", pattern=pattern) + sig.info["source"] = "builtin_msvc" + sigs.append(sig) + + # C++ exception handling + cpp_patterns = { + "__CxxFrameHandler3": b"\x48\x89\x5c\x24\x10\x48\x89\x74\x24\x18\x57", + "_CxxThrowException": b"\x48\x83\xec\x28\x48\x8b\x05", + "__CxxCallUnwindDtor": b"\x48\x83\xec\x28\x48\x8b\x01", + "_CxxFrameHandler": b"\x48\x89\x5c\x24\x10\x48\x89\x74\x24\x18", + } + for name, pattern in cpp_patterns.items(): + sig = Signature(name, library="msvc_cpp", pattern=pattern) + sig.info["source"] = "builtin_cpp" + sigs.append(sig) + + return sigs + + # ── 持久化 ── + + def save_to_store(self, lib_name: str): + """将签名库保存到 NetnodeStore""" + if not self._store: + return False + lib = self._libraries.get(lib_name) + if not lib: + return False + json_data = lib.export_json() + self._store.set_blob(f"sig_lib:{lib_name}", json_data.encode("utf-8")) + return True + + def load_from_store(self, lib_name: str) -> Optional[SignatureLibrary]: + """从 NetnodeStore 加载签名库""" + if not self._store: + return None + blob = self._store.get_blob(f"sig_lib:{lib_name}") + if not blob: + return None + json_data = blob.decode("utf-8") + lib = SignatureLibrary.import_json(json_data) + self.add_library(lib) + return lib + + # ── 扫描 ── + + def scan(self, functions: List[dict]) -> List[dict]: + """ + 扫描函数列表,匹配所有库 + + Args: + functions: [{"name": str, "address": str, "head32": bytes, "head64": bytes}, ...] + + Returns: + 匹配结果列表 + """ + self._stats["total_scans"] += 1 + all_matches = [] + + for func in functions: + func_name = func.get("name", "") + addr = func.get("address", "") + head32 = func.get("head32", b"") + head64 = func.get("head64") + + for lib_name, lib in self._libraries.items(): + sig = lib.match(head32, head64) + if sig: + all_matches.append({ + "func_name": func_name, + "address": addr, + "signature": sig, + "library": lib_name, + "matched_name": sig.name, + }) + self._stats["total_matches"] += 1 + break # 命中后不再搜索其他库 + + return all_matches + + def scan_single(self, name: str, address: str, + head32: bytes, head64: Optional[bytes] = None) -> Optional[dict]: + """扫描单个函数""" + results = self.scan([{ + "name": name, + "address": address, + "head32": head32, + "head64": head64, + }]) + return results[0] if results else None + + # ── 与 Ghidra MCP 集成 ── + + def prepare_ghidra_scan(self, functions_data: List[dict]) -> List[dict]: + """ + 将 Ghidra 函数数据转换为签名引擎扫描格式 + + Ghidra 函数数据格式: + [{"name": "FUN_1000", "address": "0x1000", "bytes": "..."}, ...] + + Returns: + 签名引擎输入格式 + """ + result = [] + for func in functions_data: + raw_bytes = func.get("bytes", "") + if isinstance(raw_bytes, str): + raw_bytes = bytes.fromhex(raw_bytes.replace(" ", "")) + head32 = raw_bytes[:32] + head64 = raw_bytes[:64] + result.append({ + "name": func.get("name", "unknown"), + "address": func.get("address", ""), + "head32": head32, + "head64": head64, + }) + return result + + # ── 统计 ── + + def get_stats(self) -> dict: + lib_stats = {} + for name, lib in self._libraries.items(): + lib_stats[name] = lib.get_stats() + + return { + **self._stats, + "libraries": lib_stats, + "hit_rate": f"{self._stats['total_matches'] / max(self._stats['total_scans'], 1):.1%}", + } + + +# ══════════════════════════════════════════════ +# Ghidra 集成工具函数 +# ══════════════════════════════════════════════ + +def signature_from_ghidra_function(func_info: dict) -> Optional[Signature]: + """ + 从 Ghidra 函数分析结果生成签名 + + Args: + func_info: Ghidra analyze_function_complete 输出 + + Returns: + 可用于匹配的 Signature 对象 + """ + # 从反编译结果提取特征字节 + disasm = func_info.get("disassembly", "") + if not disasm: + return None + + # 提取函数名 + name = func_info.get("name", func_info.get("function_name", "unknown")) + + # 尝试从 disassembly 中提取原始字节 + bytes_hex = func_info.get("bytes", "") + if isinstance(bytes_hex, str) and bytes_hex: + try: + raw = bytes.fromhex(bytes_hex.replace(" ", "").replace("\n", "")) + sig = Signature(name, library="ghidra_scan", pattern=raw) + sig.info["address"] = func_info.get("address", "") + sig.info["size"] = func_info.get("size", 0) + return sig + except (ValueError, AttributeError): + pass + + return None + + +def batch_generate_signatures(functions: List[dict]) -> SignatureLibrary: + """ + 从一批 Ghidra 函数批量生成签名库 + + Args: + functions: Ghidra list_functions 或 search_functions 的输出 + + Returns: + 包含所有函数签名的 SignatureLibrary + """ + lib = SignatureLibrary("ghidra_export") + for func in functions: + sig = signature_from_ghidra_function(func) + if sig: + lib.add(sig) + return lib