diff --git a/wechat_cli/keys/common.py b/wechat_cli/keys/common.py index c281b67..3c5921d 100644 --- a/wechat_cli/keys/common.py +++ b/wechat_cli/keys/common.py @@ -17,15 +17,35 @@ def verify_enc_key(enc_key, db_page1): - """通过 HMAC-SHA512 校验 page 1 验证 enc_key 是否正确。""" + """通过 HMAC-SHA1 (新版) 或 HMAC-SHA512 (旧版) 校验 page 1 验证 enc_key 是否正确。 + + 微信 3.9+ 使用 HMAC-SHA1,更早版本使用 HMAC-SHA512。 + """ salt = db_page1[:SALT_SZ] mac_salt = bytes(b ^ 0x3A for b in salt) - mac_key = hashlib.pbkdf2_hmac("sha512", enc_key, mac_salt, 2, dklen=KEY_SZ) - hmac_data = db_page1[SALT_SZ: PAGE_SZ - 80 + 16] - stored_hmac = db_page1[PAGE_SZ - 64: PAGE_SZ] - hm = hmac_mod.new(mac_key, hmac_data, hashlib.sha512) - hm.update(struct.pack(" len(data): + return 0 + return struct.unpack_from("= 15 and data not in {b"\x00" * KEY_SZ, b"\xff" * KEY_SZ} + + +def _find_bytes_in_regions(regions, read_fn, needle): + """在内存区域中查找字节串,返回所有匹配地址""" + addresses = set() + overlap = max(0, len(needle) - 1) + for base, size in regions: + offset = 0 + tail = b"" + tail_base = base + chunk_size = 2 * 1024 * 1024 + + while offset < size: + current_size = min(chunk_size, size - offset) + chunk = read_fn(base + offset, current_size) or b"" + data_base = tail_base if tail else base + offset + data = tail + chunk + + if data: + pos = data.find(needle) + while pos >= 0: + addresses.add(data_base + pos) + pos = data.find(needle, pos + 1) + + if overlap: + tail = data[-overlap:] + tail_base = data_base + max(0, len(data) - len(tail)) + else: + tail = b"" + tail_base = base + offset + current_size + else: + tail = b"" + tail_base = base + offset + current_size + offset += current_size + + return addresses + + +def _windows_v411_config_key_candidates(blob): + """从 Config.Cipher blob 中提取密钥候选""" + if not blob or len(blob) > WINDOWS_CONFIG_BLOB_MAX: + return [] + decoded = _xor_repeat(blob, WINDOWS_CONFIG_XOR_MASK) + out = [] + seen = set() + + for match in WINDOWS_CONFIG_LITERAL_RE.finditer(decoded): + run = match.group(1).decode("ascii").lower() + starts = [0] + if len(run) > 96: + starts.extend(range(0, len(run) - 63, 32)) + starts.append(len(run) - 64) + + for start in dict.fromkeys(starts): + if start < 0 or start + 64 > len(run): + continue + enc_key_hex = run[start:start + 64] + try: + enc_key = bytes.fromhex(enc_key_hex) + except ValueError: + continue + if not _probable_32_byte_key(enc_key): + continue + + embedded_salt = None + if start + 96 <= len(run): + embedded_salt = run[start + 64:start + 96] + + item = (enc_key_hex, embedded_salt) + if item not in seen: + seen.add(item) + out.append(item) + + return out + + +def _verify_direct_key_candidate(enc_key_hex, embedded_salt, db_files, salt_to_dbs, key_map, remaining_salts): + """验证单个密钥候选""" + if not remaining_salts: + return 0 + try: + enc_key = bytes.fromhex(enc_key_hex) + except ValueError: + return 0 + if not _probable_32_byte_key(enc_key): + return 0 + + matched = 0 + target_salts = [embedded_salt] if embedded_salt in remaining_salts else list(remaining_salts) + + for salt_hex in target_salts: + if salt_hex not in remaining_salts: + continue + for _rel, _path, _sz, s, page1 in db_files: + if s == salt_hex and verify_enc_key(enc_key, page1): + key_map[salt_hex] = enc_key_hex + remaining_salts.discard(salt_hex) + matched += 1 + break + + return matched + + +def _scan_config_cipher(pid, h, regions, db_files, salt_to_dbs, key_map, remaining_salts): + """Windows 4.1+ Config.Cipher 扫描""" + if not remaining_salts: + return 0 + + # 查找 Config.Cipher 名称字符串 + needle_addresses = _find_bytes_in_regions(regions, lambda addr, sz: _read_mem(h, addr, sz), WINDOWS_CONFIG_CIPHER_NAME) + if not needle_addresses: + return 0 + + print(f"[*] Config.Cipher 扫描: 找到 {len(needle_addresses)} 个名称匹配") + + # 构建地址+长度对的字节模式 + pair_patterns = [ + struct.pack("= 0: + qaddr = base + pos + node_base = qaddr - 0x10 + node = _read_mem(h, node_base, 0x50) + if not node or len(node) < 0x40: + pos = data.find(pattern, pos + 1) + continue + + if _u64_from(node, 0x10) not in needle_addresses or _u64_from(node, 0x18) != len(WINDOWS_CONFIG_CIPHER_NAME): + pos = data.find(pattern, pos + 1) + continue + + config_ptr = _u64_from(node, 0x28) + if not (0x10000 <= config_ptr < WINDOWS_MAX_USER_ADDRESS): + pos = data.find(pattern, pos + 1) + continue + + seen_config_ptrs.add(config_ptr) + + # 读取 Config 对象 + obj = _read_mem(h, config_ptr + 0x88, 0x28) + if not obj or len(obj) < 0x18: + pos = data.find(pattern, pos + 1) + continue + + data_ptr = _u64_from(obj, 0x8) + data_len = _u64_from(obj, 0x10) + if not (0 < data_len <= WINDOWS_CONFIG_BLOB_MAX and 0x10000 <= data_ptr < WINDOWS_MAX_USER_ADDRESS): + pos = data.find(pattern, pos + 1) + continue + + blob = _read_mem(h, data_ptr, int(data_len)) + if not blob or len(blob) != data_len: + pos = data.find(pattern, pos + 1) + continue + + # 提取并验证密钥候选 + for enc_key_hex, embedded_salt in _windows_v411_config_key_candidates(blob): + candidate = (enc_key_hex, embedded_salt) + if candidate in seen_candidates: + continue + seen_candidates.add(candidate) + candidate_count += 1 + + matched = _verify_direct_key_candidate( + enc_key_hex, embedded_salt, db_files, salt_to_dbs, + key_map, remaining_salts + ) + if matched: + matched_count += matched + + pos = data.find(pattern, pos + 1) + + if matched_count: + print(f"[+] Config.Cipher 扫描匹配 {matched_count}/{len(salt_to_dbs)} salts (候选={candidate_count})") + + return matched_count + + + def extract_keys(db_dir, output_path, pid=None): """提取 Windows 微信数据库密钥。 @@ -100,7 +323,11 @@ def extract_keys(db_dir, output_path, pid=None): all_hex_matches = 0 t0 = time.time() + # 优先尝试 Config.Cipher 扫描(Windows 4.1+) for pid_val, mem_kb in pids: + if not remaining_salts: + break + h = kernel32.OpenProcess(0x0010 | 0x0400, False, pid_val) if not h: print(f"[WARN] 无法打开进程 PID={pid_val},跳过") @@ -110,27 +337,37 @@ def extract_keys(db_dir, output_path, pid=None): regions = _enum_regions(h) total_bytes = sum(s for _, s in regions) total_mb = total_bytes / 1024 / 1024 - print(f"\n[*] 扫描 PID={pid_val} ({total_mb:.0f}MB, {len(regions)} 区域)") + print(f"\n[*] PID={pid_val} ({total_mb:.0f}MB, {len(regions)} 区域)") - scanned_bytes = 0 - for reg_idx, (base, size) in enumerate(regions): - data = _read_mem(h, base, size) - scanned_bytes += size - if not data: - continue + # 先尝试 Config.Cipher 扫描 + matched = _scan_config_cipher(pid_val, h, regions, db_files, salt_to_dbs, key_map, remaining_salts) + if matched and not remaining_salts: + print(f"[+] Config.Cipher 扫描已覆盖所有数据库") + kernel32.CloseHandle(h) + break - all_hex_matches += scan_memory_for_keys( - data, hex_re, db_files, salt_to_dbs, - key_map, remaining_salts, base, pid_val, print, - ) - - if (reg_idx + 1) % 200 == 0: - elapsed = time.time() - t0 - progress = scanned_bytes / total_bytes * 100 if total_bytes else 100 - print( - f" [{progress:.1f}%] {len(key_map)}/{len(salt_to_dbs)} salts matched, " - f"{all_hex_matches} hex patterns, {elapsed:.1f}s" + # 如果还有剩余,回退到传统扫描 + if remaining_salts: + print(f"[*] 传统内存扫描 (剩余 {len(remaining_salts)} salts)") + scanned_bytes = 0 + for reg_idx, (base, size) in enumerate(regions): + data = _read_mem(h, base, size) + scanned_bytes += size + if not data: + continue + + all_hex_matches += scan_memory_for_keys( + data, hex_re, db_files, salt_to_dbs, + key_map, remaining_salts, base, pid_val, print, ) + + if (reg_idx + 1) % 200 == 0: + elapsed = time.time() - t0 + progress = scanned_bytes / total_bytes * 100 if total_bytes else 100 + print( + f" [{progress:.1f}%] {len(key_map)}/{len(salt_to_dbs)} salts matched, " + f"{all_hex_matches} hex patterns, {elapsed:.1f}s" + ) finally: kernel32.CloseHandle(h)