Add Windows GUI and WXWork export support

This commit is contained in:
xincheng
2026-05-17 07:07:07 +08:00
38 changed files with 10005 additions and 887 deletions

View File

@@ -0,0 +1,104 @@
"""Tests for `chat_export_helpers._extract_content` group prefix handling.
Issue #88: 群聊里的引用回复 / appmsg 卡片在 export_chat / export_all_chats
渲染成 link_or_file 且 content 为空。根因是 `_extract_content` 把带
`wxid_xxx:\\n` 群前缀的原始 content 直接喂给 `_format_app_message_text`
XML 解析器在前缀文本上崩溃。
修复后:
- 检测到 chat_username 是 @chatroom先用 `_parse_message_content` 剥前缀
- 把 `is_group=True` 透传给 `_format_app_message_text` 让引用回复的发送者
标签解析走群路径
- 用真实的 contact names dict 而不是 `{}` 让 1-on-1 也能解出昵称
"""
import unittest
from unittest.mock import patch
import chat_export_helpers
import mcp_server
def _refer_appmsg(refer_content="hello world"):
"""合成一条引用回复 appmsg。"""
return (
'<msg><appmsg appid="" sdkver="0">'
'<title>quote reply</title>'
'<type>57</type>'
'<refermsg>'
'<type>1</type>'
f'<content>{refer_content}</content>'
'<fromusr>wxid_orig_sender</fromusr>'
'<displayname>Original Sender</displayname>'
'</refermsg>'
'</appmsg></msg>'
)
class ExtractContentGroupPrefixTests(unittest.TestCase):
def setUp(self):
# Skip decompression
self._patch = patch.object(
mcp_server, '_decompress_content',
side_effect=lambda content, ct: content,
)
self._patch.start()
self._names_patch = patch.object(
mcp_server, 'get_contact_names',
return_value={'wxid_orig_sender': 'Alice'},
)
self._names_patch.start()
def tearDown(self):
self._patch.stop()
self._names_patch.stop()
def test_group_appmsg_with_prefix_renders_correctly(self):
"""Issue #88: 群引用回复带 'wxid_xxx:\\n' 前缀,需要正确剥离后再解析。"""
prefixed = 'wxid_group_member:\n' + _refer_appmsg('hello group')
rendered, extras = chat_export_helpers._extract_content(
local_id=100, local_type=49, content=prefixed, ct=0,
chat_username='12345@chatroom', chat_display_name='Test Group',
)
self.assertIsNotNone(rendered, "群引用回复不应该解析失败返回 None")
self.assertIn('quote reply', rendered)
self.assertIn('hello group', rendered, "被引用内容应该出现在渲染结果里")
def test_one_on_one_appmsg_unaffected(self):
"""1-on-1 场景没有前缀,行为应该保持不变。"""
rendered, _ = chat_export_helpers._extract_content(
local_id=100, local_type=49, content=_refer_appmsg('hi'), ct=0,
chat_username='wxid_friend', chat_display_name='Friend',
)
self.assertIsNotNone(rendered)
self.assertIn('hi', rendered)
def test_group_text_prefix_stripped(self):
"""群里的 base=1 text 消息content 也带前缀,应该被剥掉。"""
text, _ = chat_export_helpers._extract_content(
local_id=100, local_type=1, content='wxid_xx:\nhello group',
ct=0, chat_username='12345@chatroom', chat_display_name='Group',
)
self.assertEqual(text, 'hello group')
def test_one_on_one_text_unaffected(self):
"""1-on-1 text 没有前缀概念,原样返回。"""
text, _ = chat_export_helpers._extract_content(
local_id=100, local_type=1, content='hello friend', ct=0,
chat_username='wxid_friend', chat_display_name='Friend',
)
self.assertEqual(text, 'hello friend')
def test_group_quote_uses_real_names(self):
"""群引用回复的发送者标签应该用真实 contact names 解析。"""
prefixed = 'wxid_group_member:\n' + _refer_appmsg()
rendered, _ = chat_export_helpers._extract_content(
local_id=100, local_type=49, content=prefixed, ct=0,
chat_username='12345@chatroom', chat_display_name='Test Group',
)
# is_group=True 走 group 分支:用 ref_user (wxid_orig_sender) 查 names
# → 'Alice'。原先 names={} 会回退到 displayname。
self.assertIn('Alice', rendered, "应该用 names dict 解析出 'Alice'")
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,97 @@
"""测试 get_chat_images 新增的 offset / start_time / end_time 参数。"""
import os
import sys
from unittest.mock import patch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import mcp_server
def _img(local_id, create_time, md5=None, size=None):
info = {'local_id': local_id, 'create_time': create_time, 'md5': md5}
if size is not None:
info['size'] = size
return info
def _run_with(shard_images_map, **kwargs):
"""Helper: stub collaborators and call get_chat_images with new kwargs."""
shards = [{'db_path': k, 'table_name': 'Msg_x'} for k in shard_images_map]
captured = {'calls': []}
def fake_list(db_path, table_name, username, limit=20, start_ts=None, end_ts=None):
captured['calls'].append({
'db_path': db_path, 'limit': limit, 'start_ts': start_ts, 'end_ts': end_ts,
})
return shard_images_map.get(db_path, [])
with patch.object(mcp_server, 'resolve_username', return_value='wxid_demo'), \
patch.object(mcp_server, 'get_contact_names', return_value={'wxid_demo': 'Demo'}), \
patch.object(mcp_server, '_find_msg_tables_for_user', return_value=shards), \
patch.object(mcp_server._image_resolver, 'list_chat_images', side_effect=fake_list):
return mcp_server.get_chat_images('Demo', **kwargs), captured
def test_invalid_offset_returns_error():
out, _ = _run_with({}, offset=-1)
assert '错误' in out
def test_invalid_time_range_returns_error():
"""start_time 晚于 end_time 应报错。"""
out, _ = _run_with({}, start_time='2026-05-10', end_time='2026-05-01')
assert '错误' in out
def test_candidate_limit_includes_offset():
"""每 shard 拉 limit+offset 张候选, 保证全局分页能切到正确的页。"""
_, captured = _run_with({'/a': [_img(1, 1000)]}, limit=5, offset=10)
assert captured['calls'][0]['limit'] == 15
def test_start_end_ts_forwarded_to_shard_query():
"""start_time / end_time 解析为 unix 秒后透传给 shard 查询。"""
_, captured = _run_with(
{'/a': []},
start_time='2026-05-01',
end_time='2026-05-31',
)
call = captured['calls'][0]
assert call['start_ts'] is not None
assert call['end_ts'] is not None
assert call['start_ts'] < call['end_ts']
def test_offset_slices_paged_window():
"""offset=2, limit=2 取全局排序后第 3-4 张图片。"""
shard_a = [_img(1, 1100), _img(2, 1000)]
shard_b = [_img(3, 1300), _img(4, 1200)]
out, _ = _run_with({'/a': shard_a, '/b': shard_b}, limit=2, offset=2)
# 全局排序后顺序: 1300, 1200, 1100, 1000 → 第 3-4 是 1100, 1000 → local_id 1, 2
assert 'local_id=1' in out
assert 'local_id=2' in out
assert 'local_id=3' not in out
assert 'local_id=4' not in out
def test_header_shows_time_range_when_given():
shard_a = [_img(1, 1000, md5='abc')]
out, _ = _run_with({'/a': shard_a}, start_time='2026-05-01')
assert '时间范围' in out
assert '2026-05-01' in out
def test_header_shows_offset_limit():
shard_a = [_img(1, 1000, md5='abc')]
out, _ = _run_with({'/a': shard_a}, limit=10, offset=20)
assert 'offset=20' in out
assert 'limit=10' in out
def test_default_behavior_unchanged():
"""不传新参数时行为与旧接口一致 — offset=0 切片就是 [:limit]。"""
shard_a = [_img(1, 1100, md5='a1'), _img(2, 1000, md5='a2')]
out, _ = _run_with({'/a': shard_a})
assert 'local_id=1' in out
assert 'local_id=2' in out

View File

@@ -0,0 +1,481 @@
"""ImageResolver 在 V2 加密格式下的端到端解密测试。
覆盖:
- v2_decrypt_file 能正确还原 AES-ECB + XOR 混合加密的合成数据
- decrypt_dat_file 按 magic 自动分发 V2 / V1 / 老 XOR 三条路径
- ImageResolver 通过 __init__ 注入 aes_key/xor_key 后,能端到端解密 V2 .dat
- 没传 aes_key 时遇到 V2 文件返回结构化错误,而不是 crash 或返回错误数据
- 默认参数下老 XOR 路径不受影响,保持向后兼容
"""
import hashlib
import os
import sqlite3
import struct
import tempfile
import unittest
from Crypto.Cipher import AES
from Crypto.Util import Padding
from decode_image import (
V1_MAGIC_FULL,
V2_MAGIC_FULL,
ImageResolver,
decrypt_dat_file,
v2_decrypt_file,
)
# 测试用 16 字节 AES key (任意值,仅用于合成测试数据)
TEST_AES_KEY = b'1234567890abcdef'
TEST_XOR_KEY = 0x37
# 最小可识别的 PNG payload (含 IHDR 和 IEND chunk),长度 88 字节
TEST_PNG_PAYLOAD = (
b'\x89PNG\r\n\x1a\n'
+ b'\x00\x00\x00\rIHDR'
+ b'\x00' * 64
+ b'IEND\xaeB`\x82'
)
def _build_v2_dat(plaintext, aes_size, xor_size,
aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY,
magic=V2_MAGIC_FULL):
"""构造合成的 V2 / V1 .dat 字节串。
布局: [6B magic][4B aes_size LE][4B xor_size LE][1B pad][AES-ECB][raw][XOR]
aes_size / xor_size 是明文字段长度,AES 段做 PKCS7 padding 后向上对齐到 16 倍数。
"""
if aes_size + xor_size > len(plaintext):
raise ValueError("aes_size + xor_size 超过 plaintext 长度")
aes_plain = plaintext[:aes_size]
raw_plain = plaintext[aes_size:len(plaintext) - xor_size]
xor_plain = plaintext[len(plaintext) - xor_size:]
cipher = AES.new(aes_key[:16], AES.MODE_ECB)
aes_cipher = cipher.encrypt(Padding.pad(aes_plain, AES.block_size))
xor_cipher = bytes(b ^ xor_key for b in xor_plain)
header = magic + struct.pack('<LL', aes_size, xor_size) + b'\x00'
return header + aes_cipher + raw_plain + xor_cipher
class _FakeCache:
"""ImageResolver 测试用最小缓存桩,绕过真实 DB 解密。"""
def __init__(self, mapping):
self._mapping = mapping
def get(self, rel_key):
return self._mapping.get(rel_key)
def _make_resource_db(path, local_id, file_md5, username="wxid_test123",
chat_id=1, message_create_time=1700000000,
message_local_type=3, extra_rows=()):
"""构造最小 message_resource.db, 表 schema 对齐真实微信结构。
真实表里 message_local_id 不全局唯一 (跨 chat 重复, 活跃 chat 内也会复用),
解析必须用 ChatName2Id.rowid -> chat_id 限定 + message_local_type=3 过滤图片。
packed_info 里嵌入 extract_md5_from_packed_info 期望的 protobuf marker
(\\x12\\x22\\x0a\\x20) 加 32 字节 ASCII hex MD5。
Args:
extra_rows: 额外 (chat_id, message_local_id, message_local_type,
message_create_time, file_md5) 元组列表, 用于构造同 local_id
跨 chat / 同 chat 多版本的歧义场景。
"""
marker = b'\x12\x22\x0a\x20'
def _packed(md5_hex):
return b'\x00' * 8 + marker + md5_hex.encode('ascii') + b'\x00' * 4
conn = sqlite3.connect(path)
try:
conn.execute("""
CREATE TABLE MessageResourceInfo (
message_id INTEGER PRIMARY KEY,
chat_id INTEGER,
sender_id INTEGER,
message_local_type INTEGER,
message_create_time INTEGER,
message_local_id INTEGER,
message_svr_id INTEGER,
message_origin_source INTEGER,
packed_info BLOB
)
""")
conn.execute(
"CREATE TABLE ChatName2Id (user_name TEXT PRIMARY KEY, update_time INTEGER)"
)
conn.execute(
"INSERT INTO ChatName2Id (rowid, user_name, update_time) VALUES (?, ?, ?)",
(chat_id, username, message_create_time),
)
next_msg_id = 1
conn.execute(
"INSERT INTO MessageResourceInfo "
"(message_id, chat_id, sender_id, message_local_type, message_create_time, "
" message_local_id, message_svr_id, message_origin_source, packed_info) "
"VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
(next_msg_id, chat_id, 0, message_local_type, message_create_time,
local_id, 0, 0, _packed(file_md5)),
)
next_msg_id += 1
for extra in extra_rows:
ex_chat_id, ex_local_id, ex_type, ex_ctime, ex_md5 = extra
conn.execute(
"INSERT INTO MessageResourceInfo "
"(message_id, chat_id, sender_id, message_local_type, message_create_time, "
" message_local_id, message_svr_id, message_origin_source, packed_info) "
"VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
(next_msg_id, ex_chat_id, 0, ex_type, ex_ctime, ex_local_id, 0, 0, _packed(ex_md5)),
)
next_msg_id += 1
conn.commit()
finally:
conn.close()
class TestV2DecryptSynthetic(unittest.TestCase):
"""v2_decrypt_file / decrypt_dat_file 在合成数据上的正确性"""
def test_v2_round_trip_recovers_payload(self):
# 合成 V2 .dat 解密后字节级等于原始 payload
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
out_path, fmt = v2_decrypt_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
with open(out_path, 'rb') as f:
self.assertEqual(f.read(), TEST_PNG_PAYLOAD)
def test_decrypt_dat_file_routes_v2_by_magic(self):
# decrypt_dat_file 看到 V2 magic 应自动走 V2 路径
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
out_path, fmt = decrypt_dat_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
def test_decrypt_dat_file_v1_uses_fixed_key(self):
# V1 magic 走固定 key,即便外部不传 aes_key 也能解密
v1_fixed_key = b'cfcd208495d565ef' # md5("0")[:16],由 v2_decrypt_file 内部使用
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(
TEST_PNG_PAYLOAD, aes_size=32, xor_size=16,
aes_key=v1_fixed_key, magic=V1_MAGIC_FULL,
))
out_path, fmt = decrypt_dat_file(
dat_path, aes_key=None, xor_key=TEST_XOR_KEY
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
def test_decrypt_dat_file_legacy_xor_route(self):
# 老 XOR 格式 (无 V1/V2 magic),decrypt_dat_file 应回退到 xor_decrypt_file 不需要 aes_key
xor_key = 0x37
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(bytes(b ^ xor_key for b in TEST_PNG_PAYLOAD))
out_path, fmt = decrypt_dat_file(dat_path, aes_key=None)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
def test_v2_accepts_str_aes_key_from_config(self):
# 真实场景下 aes_key 来自 config.json,是 ASCII string 不是 bytes;
# v2_decrypt_file 内部应自行 encode,避免 TypeError
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
out_path, fmt = decrypt_dat_file(
dat_path, aes_key=TEST_AES_KEY.decode('ascii'), xor_key=TEST_XOR_KEY,
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
def test_v2_accepts_str_xor_key_from_config(self):
# 与 aes_key 的 str 处理对称: config.json 里把 xor_key 写成 "0x88" / "136" 也应能正常解密
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
out_path, fmt = decrypt_dat_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=hex(TEST_XOR_KEY),
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
def test_v2_wxgf_payload_returns_hevc_format(self):
# 微信 V2 动图 (wxgf 裸流 HEVC) 解密后 fmt='hevc',输出文件以 .hevc 结尾;
# 当前 ImageResolver 不再向 JPEG 转 (那是 monitor_web 的职责),保持原样输出。
wxgf_payload = b'wxgf' + b'\x00' * 84 # 88 字节,与 PNG payload 同长度,避免改 aes/xor sizes
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(wxgf_payload, aes_size=32, xor_size=16))
out_path, fmt = decrypt_dat_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY,
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'hevc')
self.assertTrue(out_path.endswith('.hevc'))
def test_v2_rejects_wrong_aes_key(self):
# AES key 错时 detect_image_format 返回 'bin' (magic 不识别),v2_decrypt_file
# 应拒绝写出 .bin 垃圾文件并返回 (None, None),让 caller 知道解密失败。
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
wrong_aes_key = b'wrongkey00000000'
out_path, fmt = v2_decrypt_file(
dat_path, aes_key=wrong_aes_key, xor_key=TEST_XOR_KEY,
)
self.assertIsNone(out_path)
self.assertIsNone(fmt)
def test_v2_rejects_wrong_xor_key_jpg_trailer(self):
# JPG 必须以 FF D9 (EOI) 收尾。XOR key 错时尾部 16 字节乱码,
# FF D9 被破坏,触发尾部 magic 校验失败。
jpg_payload = b'\xff\xd8\xff' + b'\x00' * 83 + b'\xff\xd9' # 88 bytes, FF D9 在末尾
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(jpg_payload, aes_size=32, xor_size=16))
# 翻转所有 XOR 字节: TEST_XOR_KEY ^ 0xff 保证每字节都错位
out_path, fmt = v2_decrypt_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY ^ 0xff,
)
self.assertIsNone(out_path)
self.assertIsNone(fmt)
def test_v2_rejects_wrong_xor_key_png_iend(self):
# PNG 末尾 12 字节必须含 IEND chunk。XOR key 错时 IEND 被破坏。
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
out_path, fmt = v2_decrypt_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY ^ 0xff,
)
self.assertIsNone(out_path)
self.assertIsNone(fmt)
def test_v2_skip_xor_validation_when_xor_size_zero(self):
# xor_size < 2 时没有 XOR 段(或样本不足以验证),不应触发尾部 magic 校验。
# 构造 xor_size=0 的 PNG (整张图都在 AES + raw 段),xor_key 实际不参与解密。
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=0))
# xor_key 传 0 也应成功 (XOR 段长度 0)
out_path, fmt = v2_decrypt_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=0x00,
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'png')
def test_v2_wxgf_skips_trailer_validation(self):
# wxgf (HEVC 裸流) 没有强制 trailer signature,XOR key 错时也不应被尾部校验误杀
# (wxgf 路径在 elif 链前面命中,直接 fmt='hevc',不进入 XOR 校验分支)。
# 这里验证:即便 XOR key 错导致末尾字节乱码,只要 wxgf magic 在头部正确,
# 仍按 hevc 输出 — 因为我们只校验 jpg/png,其他格式跳过。
wxgf_payload = b'wxgf' + b'\x00' * 84
with tempfile.TemporaryDirectory() as td:
dat_path = os.path.join(td, "test.dat")
with open(dat_path, 'wb') as f:
f.write(_build_v2_dat(wxgf_payload, aes_size=32, xor_size=16))
out_path, fmt = v2_decrypt_file(
dat_path, aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY ^ 0xff,
)
self.assertIsNotNone(out_path)
self.assertEqual(fmt, 'hevc')
class TestImageResolverV2(unittest.TestCase):
"""ImageResolver 端到端:从 local_id 到解密文件,验证 V2 keys 注入路径"""
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
tmp = self._tmp.name
self.wechat_base = os.path.join(tmp, "wechat")
self.out_dir = os.path.join(tmp, "decoded")
os.makedirs(self.out_dir, exist_ok=True)
self.username = "wxid_test123"
self.local_id = 42
self.file_md5 = "0123456789abcdef0123456789abcdef"
username_hash = hashlib.md5(self.username.encode()).hexdigest()
img_dir = os.path.join(
self.wechat_base, "msg", "attach", username_hash, "2025-08", "Img"
)
os.makedirs(img_dir, exist_ok=True)
self.dat_path = os.path.join(img_dir, f"{self.file_md5}.dat")
with open(self.dat_path, 'wb') as f:
f.write(_build_v2_dat(TEST_PNG_PAYLOAD, aes_size=32, xor_size=16))
self.db_path = os.path.join(tmp, "message_resource.db")
_make_resource_db(self.db_path, self.local_id, self.file_md5)
self.cache = _FakeCache({"message/message_resource.db": self.db_path})
def test_decode_image_v2_with_keys(self):
resolver = ImageResolver(
self.wechat_base, self.out_dir, self.cache,
aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY,
)
result = resolver.decode_image(self.username, self.local_id)
self.assertTrue(result['success'], msg=result)
self.assertEqual(result['format'], 'png')
self.assertEqual(result['md5'], self.file_md5)
with open(result['path'], 'rb') as f:
self.assertEqual(f.read(), TEST_PNG_PAYLOAD)
def test_decode_image_v2_missing_aes_key_returns_error(self):
# 没传 aes_key 时遇到 V2 文件应返回 success=False,而不是 crash 或写入错误文件
resolver = ImageResolver(
self.wechat_base, self.out_dir, self.cache, aes_key=None,
)
result = resolver.decode_image(self.username, self.local_id)
self.assertFalse(result['success'])
self.assertIn('AES key', result['error'])
self.assertEqual(result['md5'], self.file_md5)
def test_decode_image_default_args_preserve_legacy_xor(self):
# 默认参数 (aes_key=None) + 老 XOR .dat 应保持向后兼容
os.unlink(self.dat_path)
legacy_xor_key = 0x37
with open(self.dat_path, 'wb') as f:
f.write(bytes(b ^ legacy_xor_key for b in TEST_PNG_PAYLOAD))
resolver = ImageResolver(self.wechat_base, self.out_dir, self.cache)
result = resolver.decode_image(self.username, self.local_id)
self.assertTrue(result['success'], msg=result)
self.assertEqual(result['format'], 'png')
def test_decode_image_disambiguates_local_id_across_chats(self):
"""同 local_id 跨 chat 重复时, 必须按 username -> chat_id 选对; 否则会拿到
别的 chat 的 MD5 (或视频 type=43 的 packed_info), 解出错图。
生产 DB 上同一个 message_local_id 实测会出现在 5+ 个不同 chat 里,
其中混有图片 (type=3) / 视频 (type=43) / 群聊 / 私聊, 必须 chat-scoped
+ type 过滤才能定位。
"""
os.unlink(self.db_path)
other_md5 = "f" * 32
video_md5 = "a" * 32
_make_resource_db(
self.db_path, self.local_id, self.file_md5,
username=self.username, chat_id=7,
message_create_time=1778487726,
extra_rows=[
# 另一个 chat 同 local_id 同图片类型, MD5 不同 —— 选错就拿这个
(5, self.local_id, 3, 1700000000, other_md5),
# 又一个 chat 同 local_id 但是视频 (type=43), 应被 type 过滤
(132, self.local_id, 43, 1750000000, video_md5),
],
)
# 给冲突 chat 也注册 user_name, 否则 chat-scope 等价
conn = sqlite3.connect(self.db_path)
conn.execute(
"INSERT INTO ChatName2Id (rowid, user_name, update_time) VALUES (5, ?, 0), (132, ?, 0)",
("other_chat_wxid", "video_chat_wxid"),
)
conn.commit()
conn.close()
resolver = ImageResolver(
self.wechat_base, self.out_dir, self.cache,
aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY,
)
result = resolver.decode_image(self.username, self.local_id)
self.assertTrue(result['success'], msg=result)
# 必须拿目标 chat 的 MD5, 不是 other_chat 也不是视频
self.assertEqual(result['md5'], self.file_md5)
def test_decode_image_picks_latest_when_same_chat_local_id_reused(self):
"""活跃 chat 里 local_id 会被复用 (实测同 chat 同 local_id 最多 7 条);
默认应返回 message_create_time 最新的那张, 对应用户最近一次 reference。
"""
os.unlink(self.db_path)
old_md5 = "c" * 32
# self.file_md5 / self.local_id 在 _make_resource_db 默认插入为 "latest" 那条
_make_resource_db(
self.db_path, self.local_id, self.file_md5,
username=self.username, chat_id=1,
message_create_time=1778487726,
extra_rows=[
# 同 chat 同 local_id 但更早, 不应该被选中
(1, self.local_id, 3, 1700000000, old_md5),
],
)
resolver = ImageResolver(
self.wechat_base, self.out_dir, self.cache,
aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY,
)
result = resolver.decode_image(self.username, self.local_id)
self.assertTrue(result['success'], msg=result)
self.assertEqual(result['md5'], self.file_md5)
def test_decode_image_unknown_chat_returns_error(self):
"""username 在 ChatName2Id 里找不到时, 应返回结构化错误而不是 crash 或乱选 row。"""
resolver = ImageResolver(
self.wechat_base, self.out_dir, self.cache,
aes_key=TEST_AES_KEY, xor_key=TEST_XOR_KEY,
)
result = resolver.decode_image("wxid_does_not_exist", self.local_id)
self.assertFalse(result['success'])
self.assertIn('wxid_does_not_exist', result['error'])
def test_decode_image_v1_no_aes_key_uses_fixed_key(self):
# V1 magic 不会被 is_v2_format guard 拦截 (V1 magic 是 \x07\x08V1, V2 是 \x07\x08V2);
# 即便 ImageResolver(aes_key=None), V1 文件也应通过 decrypt_dat_file 内置固定 key 解密
os.unlink(self.dat_path)
v1_fixed_key = b'cfcd208495d565ef'
with open(self.dat_path, 'wb') as f:
f.write(_build_v2_dat(
TEST_PNG_PAYLOAD, aes_size=32, xor_size=16,
aes_key=v1_fixed_key, magic=V1_MAGIC_FULL,
))
# xor_key 必须跟 _build_v2_dat 加密时用的一致,否则 XOR 段乱码,
# 触发新的尾部 magic 校验失败 (PNG IEND chunk 错位)。
resolver = ImageResolver(
self.wechat_base, self.out_dir, self.cache,
aes_key=None, xor_key=TEST_XOR_KEY,
)
result = resolver.decode_image(self.username, self.local_id)
self.assertTrue(result['success'], msg=result)
self.assertEqual(result['format'], 'png')
if __name__ == '__main__':
unittest.main()

View File

@@ -0,0 +1,295 @@
"""decode_image.decode_all_dats() batch CLI 行为测试。
覆盖:
- 路径扫描:glob 命中 attach/<chat_hash>/<YYYY-MM>/Img/*.dat
- 路径解析:chat_hash / YYYY-MM 提取,_t / _h 后缀移除归并到原图 basename
- 幂等性:目标 basename 已存在(任何扩展名)时跳过;--force 强制重解
- 原子写:写到 tmp 再 os.replace;失败/异常路径不留 .tmp
- V2 无 key:计入 skipped_no_key 而非 failed
- 错误隔离:单文件异常不阻塞批次;返回失败计数
decrypt_dat_file 用 mock 隔离(避免依赖真实加密图片);is_v2_format
单独覆盖真实 magic 检测路径。
"""
import os
import struct
import tempfile
import unittest
from contextlib import redirect_stderr
import io
from unittest.mock import patch
import decode_image
def _write(path, data):
os.makedirs(os.path.dirname(path), exist_ok=True)
with open(path, "wb") as f:
f.write(data)
def _v2_magic_bytes():
# 仅用于让 is_v2_format() 返回 True
return decode_image.V2_MAGIC_FULL + struct.pack("<LL", 0, 0) + b"\x00"
def _v1_magic_bytes():
return decode_image.V1_MAGIC_FULL + struct.pack("<LL", 0, 0) + b"\x00"
class _MockedDecrypt:
"""mock decrypt_dat_file:不真解密,只往 tmp 写一个 marker 字节串然后返回 (tmp, ext)。
通过实例化时配置返回的 ext / 是否抛异常 / 是否返回 (None, None),覆盖
各种成功/失败路径。
"""
def __init__(self, ext="jpg", marker=b"DECODED", returns_none=False, raises=None):
self.ext = ext
self.marker = marker
self.returns_none = returns_none
self.raises = raises
self.calls = []
def __call__(self, dat_path, out_path=None, aes_key=None, xor_key=0x88):
self.calls.append((dat_path, out_path, aes_key, xor_key))
if self.raises:
raise self.raises
if self.returns_none:
return None, None
# 写 tmp(decode_all_dats 期望我们写完才能 os.replace)
os.makedirs(os.path.dirname(out_path), exist_ok=True)
with open(out_path, "wb") as f:
f.write(self.marker)
return out_path, self.ext
def _make_dat(attach_dir, chat_hash, ym, basename, content=None):
"""在 attach_dir 下造一个 .dat 文件,返回完整路径。content 默认是非 V2 占位。"""
if content is None:
content = b"\x00\x00\x00\x00" # 非 V2 / 非 V1 magic
p = os.path.join(attach_dir, chat_hash, ym, "Img", f"{basename}.dat")
_write(p, content)
return p
class PathParsingTests(unittest.TestCase):
"""路径扫描 / 解析 / _t _h 归并。"""
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
self.attach = os.path.join(self._tmp.name, "attach")
self.out = os.path.join(self._tmp.name, "out")
def test_finds_dat_files_under_chat_month_img(self):
_make_dat(self.attach, "hash1", "2026-01", "abc123")
_make_dat(self.attach, "hash2", "2026-02", "def456")
mock = _MockedDecrypt()
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
self.assertEqual(stats["total"], 2)
self.assertEqual(stats["decoded"], 2)
self.assertEqual(stats["failed"], 0)
def test_strips_t_suffix(self):
_make_dat(self.attach, "hash1", "2026-01", "abc123_t")
mock = _MockedDecrypt(ext="png")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
# 期望产出: out/hash1/2026-01/abc123.png(_t 已被剥)
produced = os.path.join(self.out, "hash1", "2026-01", "abc123.png")
self.assertTrue(os.path.exists(produced), f"missing: {produced}")
def test_strips_h_suffix(self):
_make_dat(self.attach, "hash1", "2026-01", "abc_h")
mock = _MockedDecrypt(ext="jpg")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
produced = os.path.join(self.out, "hash1", "2026-01", "abc.jpg")
self.assertTrue(os.path.exists(produced))
def test_mirrors_chat_and_month(self):
_make_dat(self.attach, "abcdef0123456789", "2026-04", "img1")
mock = _MockedDecrypt(ext="jpg")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
produced = os.path.join(self.out, "abcdef0123456789", "2026-04", "img1.jpg")
self.assertTrue(os.path.exists(produced))
class IdempotentTests(unittest.TestCase):
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
self.attach = os.path.join(self._tmp.name, "attach")
self.out = os.path.join(self._tmp.name, "out")
def test_existing_target_basename_skipped(self):
_make_dat(self.attach, "hash1", "2026-01", "img1")
# 预先放一个目标(任何扩展名)
existing = os.path.join(self.out, "hash1", "2026-01", "img1.png")
_write(existing, b"OLD")
mock = _MockedDecrypt()
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
self.assertEqual(stats["skipped"], 1)
self.assertEqual(stats["decoded"], 0)
self.assertEqual(len(mock.calls), 0, "decrypt_dat_file 不该被调用")
# 目标内容未被改写
with open(existing, "rb") as f:
self.assertEqual(f.read(), b"OLD")
def test_force_overrides_skip(self):
_make_dat(self.attach, "hash1", "2026-01", "img1")
existing = os.path.join(self.out, "hash1", "2026-01", "img1.png")
_write(existing, b"OLD")
mock = _MockedDecrypt(ext="jpg", marker=b"NEW")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16,
force=True, progress_every=None,
)
self.assertEqual(stats["decoded"], 1)
self.assertEqual(stats["skipped"], 0)
# 新文件以新 ext 落盘
new_file = os.path.join(self.out, "hash1", "2026-01", "img1.jpg")
self.assertTrue(os.path.exists(new_file))
def test_skip_ignores_tmp_files(self):
"""残留的 .tmp 不应该被当成"已存在目标"误判跳过。"""
_make_dat(self.attach, "hash1", "2026-01", "img1")
# 模拟之前一次中断留下的 .tmp
leftover_tmp = os.path.join(self.out, "hash1", "2026-01", "img1.unknown.tmp")
_write(leftover_tmp, b"PARTIAL")
mock = _MockedDecrypt(ext="jpg")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
self.assertEqual(stats["decoded"], 1, "残留 .tmp 不应该阻止重解")
class AtomicWriteTests(unittest.TestCase):
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
self.attach = os.path.join(self._tmp.name, "attach")
self.out = os.path.join(self._tmp.name, "out")
_make_dat(self.attach, "hash1", "2026-01", "img1")
def test_success_path_no_tmp_leftover(self):
mock = _MockedDecrypt(ext="jpg")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
target_dir = os.path.join(self.out, "hash1", "2026-01")
leftovers = [f for f in os.listdir(target_dir) if f.endswith(".tmp")]
self.assertEqual(leftovers, [], "成功路径不应有 .tmp 残留")
def test_decrypt_returns_none_no_tmp_leftover(self):
mock = _MockedDecrypt(returns_none=True)
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
self.assertEqual(stats["failed"], 1)
# decrypt 没写 tmp(returns_none=True 时也不写),所以目录可能不存在或为空
target_dir = os.path.join(self.out, "hash1", "2026-01")
if os.path.isdir(target_dir):
leftovers = [f for f in os.listdir(target_dir) if f.endswith(".tmp")]
self.assertEqual(leftovers, [])
def test_decrypt_raises_no_tmp_leftover(self):
mock = _MockedDecrypt(raises=RuntimeError("synthetic decrypt failure"))
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key="x" * 16, progress_every=None,
)
self.assertEqual(stats["failed"], 1)
target_dir = os.path.join(self.out, "hash1", "2026-01")
if os.path.isdir(target_dir):
leftovers = [f for f in os.listdir(target_dir) if f.endswith(".tmp")]
self.assertEqual(leftovers, [], "异常路径必须清理 .tmp")
class V2NoKeyTests(unittest.TestCase):
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
self.attach = os.path.join(self._tmp.name, "attach")
self.out = os.path.join(self._tmp.name, "out")
def test_v2_dat_with_no_aes_key_skipped(self):
_make_dat(self.attach, "hash1", "2026-01", "v2img", content=_v2_magic_bytes())
mock = _MockedDecrypt()
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key=None, progress_every=None,
)
self.assertEqual(stats["skipped_no_key"], 1)
self.assertEqual(stats["decoded"], 0)
self.assertEqual(stats["failed"], 0)
self.assertEqual(len(mock.calls), 0, "无 key 的 V2 文件不应该走 decrypt_dat_file")
def test_v1_dat_with_no_aes_key_still_decoded(self):
"""V1 用固定 AES key,不需要 image_aes_key,仍应被处理。"""
_make_dat(self.attach, "hash1", "2026-01", "v1img", content=_v1_magic_bytes())
mock = _MockedDecrypt(ext="jpg")
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
stats = decode_image.decode_all_dats(
self.attach, self.out, aes_key=None, progress_every=None,
)
# is_v2_format 只识别 V2(纯 V2 magic),V1 不算 V2,所以会进入 decrypt 流程
self.assertEqual(stats["decoded"], 1)
self.assertEqual(stats["skipped_no_key"], 0)
class CallbackTests(unittest.TestCase):
def test_on_file_callback_fires_per_file(self):
with tempfile.TemporaryDirectory() as tmp:
attach = os.path.join(tmp, "attach")
out = os.path.join(tmp, "out")
_make_dat(attach, "h1", "2026-01", "a")
_make_dat(attach, "h1", "2026-01", "b")
_make_dat(attach, "h1", "2026-01", "c")
events = []
mock = _MockedDecrypt()
with patch.object(decode_image, "decrypt_dat_file", mock), \
redirect_stderr(io.StringIO()):
decode_image.decode_all_dats(
attach, out, aes_key="x" * 16, progress_every=None,
on_file=lambda i, total, p, status, fmt: events.append(status),
)
self.assertEqual(len(events), 3)
self.assertTrue(all(s == "decoded" for s in events))
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,770 @@
"""单元测试find_image_key_macos 派生算法 + 端到端 smoke。
不依赖真实微信数据;用 tempdir + 合成密文构造测试。
"""
import hashlib
import json
import multiprocessing
import os
import queue as _queue_mod
import tempfile
import unittest
from unittest.mock import patch
from Crypto.Cipher import AES
import find_image_key_macos as fkm
class NormalizeWxidTests(unittest.TestCase):
def test_wxid_with_extra_segments_keeps_only_first(self):
# wxid_<seg> 形式只保留第一段下划线之内的内容
self.assertEqual(fkm.normalize_wxid("wxid_abc123_extra_more"), "wxid_abc123")
def test_wxid_no_extra_segments(self):
self.assertEqual(fkm.normalize_wxid("wxid_abc123"), "wxid_abc123")
def test_account_with_4char_alnum_suffix_stripped(self):
# macOS 路径常见your_wxid_a1b2 → your_wxid
self.assertEqual(fkm.normalize_wxid("your_wxid_a1b2"), "your_wxid")
def test_account_without_recognizable_suffix_returned_asis(self):
self.assertEqual(fkm.normalize_wxid("simple"), "simple")
self.assertEqual(fkm.normalize_wxid("foo_bar_baz"), "foo_bar_baz") # baz 是 3 char
def test_empty_or_none_returns_empty(self):
self.assertEqual(fkm.normalize_wxid(""), "")
self.assertEqual(fkm.normalize_wxid(None), "")
self.assertEqual(fkm.normalize_wxid(" "), "")
class DeriveImageKeysTests(unittest.TestCase):
def test_xor_is_low_byte_of_code(self):
xor, _ = fkm.derive_image_keys(0x12345678, "anything")
self.assertEqual(xor, 0x78)
def test_xor_handles_small_codes(self):
self.assertEqual(fkm.derive_image_keys(0xFF, "x")[0], 0xFF)
self.assertEqual(fkm.derive_image_keys(0x00, "x")[0], 0x00)
def test_aes_is_md5_hex_truncated_to_16(self):
# Golden value: 合成 fixture (uin=12345678) 派生; 算法正确性由公式
# md5(str(uin)+wxid)[:16] 决定, 测试值无需对应任何真实账号。
xor, aes = fkm.derive_image_keys(12345678, "your_wxid")
self.assertEqual(xor, 0x4E) # 12345678 & 0xFF
self.assertEqual(aes, "a0c093edddc98490")
def test_aes_does_not_normalize_wxid_internally(self):
# 归一化由调用方负责;不同 wxid 字符串产出不同 key
_, aes_full = fkm.derive_image_keys(12345678, "your_wxid_a1b2")
_, aes_norm = fkm.derive_image_keys(12345678, "your_wxid")
self.assertNotEqual(aes_full, aes_norm)
class DeriveKvcommDirCandidatesTests(unittest.TestCase):
def test_canonical_macos_path_is_first_candidate(self):
db_dir = (
"/Users/x/Library/Containers/com.tencent.xinWeChat/Data/Documents/"
"xwechat_files/wxid_abc/db_storage"
)
candidates = fkm.derive_kvcomm_dir_candidates(db_dir)
self.assertGreater(len(candidates), 0)
expected_primary = (
"/Users/x/Library/Containers/com.tencent.xinWeChat/Data/Documents/"
"app_data/net/kvcomm"
)
self.assertEqual(candidates[0], expected_primary)
def test_returns_multiple_candidates(self):
# 多候选是 Round 1 review 的关键修复点:跨版本路径覆盖
db_dir = (
"/Users/x/Library/Containers/com.tencent.xinWeChat/Data/Documents/"
"xwechat_files/wxid_abc/db_storage"
)
candidates = fkm.derive_kvcomm_dir_candidates(db_dir)
self.assertGreaterEqual(len(candidates), 3,
"应返回多个候选路径以覆盖不同微信版本布局")
def test_no_xwechat_files_still_returns_home_fallback(self):
# 即使无法从 db_dir 推算,也至少返回 HOME 默认路径作兜底
candidates = fkm.derive_kvcomm_dir_candidates("/random/path")
self.assertGreaterEqual(len(candidates), 1)
self.assertTrue(any("Containers/com.tencent.xinWeChat" in c
for c in candidates))
def test_candidates_are_unique(self):
db_dir = "/x/y/Documents/xwechat_files/wxid_abc/db_storage"
candidates = fkm.derive_kvcomm_dir_candidates(db_dir)
self.assertEqual(len(candidates), len(set(candidates)))
class FindExistingKvcommDirTests(unittest.TestCase):
def test_returns_first_existing_candidate(self):
with tempfile.TemporaryDirectory() as tmp:
# 构造合法 db_dir 路径,在第一个候选位置创建实际目录
base = os.path.join(tmp, "Documents", "xwechat_files", "wxid_x")
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
kvcomm = os.path.join(tmp, "Documents", "app_data", "net", "kvcomm")
os.makedirs(kvcomm)
self.assertEqual(fkm.find_existing_kvcomm_dir(db_dir), kvcomm)
def test_returns_none_when_no_candidate_exists(self):
# 即使 HOME fallback 候选也不存在时,应返回 None。
# 隔离测试不能依赖宿主机有/无微信安装patch expanduser 指向 tmp。
with tempfile.TemporaryDirectory() as fake_home:
with patch("os.path.expanduser", return_value=fake_home):
self.assertIsNone(fkm.find_existing_kvcomm_dir("/nonexistent/x/y/z"))
class CollectKvcommCodesTests(unittest.TestCase):
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
self.kvdir = self._tmp.name
def _touch(self, name):
with open(os.path.join(self.kvdir, name), "w") as f:
f.write("")
def test_extracts_code_from_filename(self):
# 长格式: 模拟真实 kvcomm 缓存文件命名 (合成 ID/时间戳, 测 regex 提
# uin 的能力, 不绑定任何真实账号)
self._touch("key_12345678_1111111111_1_1700000000_22222_3600_input.statistic")
self._touch("key_99999999_yyy_zzz.statistic")
self.assertEqual(fkm.collect_kvcomm_codes(self.kvdir), [12345678, 99999999])
def test_ignores_files_with_non_numeric_first_segment(self):
self._touch("key_reportnow_12345678_xxx.statistic")
self._touch("key_abc_def.statistic")
self._touch("config.ini")
self._touch("monitordata_x")
self.assertEqual(fkm.collect_kvcomm_codes(self.kvdir), [])
def test_dedupes_same_code_across_files(self):
self._touch("key_42_a.statistic")
self._touch("key_42_b.statistic")
self.assertEqual(fkm.collect_kvcomm_codes(self.kvdir), [42])
def test_missing_dir_returns_empty(self):
self.assertEqual(fkm.collect_kvcomm_codes("/nonexistent/xxx"), [])
def test_none_dir_returns_empty(self):
self.assertEqual(fkm.collect_kvcomm_codes(None), [])
class CollectWxidCandidatesTests(unittest.TestCase):
def test_returns_raw_and_normalized_when_different(self):
db_dir = "/x/Documents/xwechat_files/your_wxid_a1b2/db_storage"
self.assertEqual(fkm.collect_wxid_candidates(db_dir),
["your_wxid_a1b2", "your_wxid"])
def test_returns_one_when_normalize_is_identity(self):
db_dir = "/x/Documents/xwechat_files/wxid_abc/db_storage"
self.assertEqual(fkm.collect_wxid_candidates(db_dir), ["wxid_abc"])
def test_no_xwechat_files_returns_empty(self):
self.assertEqual(fkm.collect_wxid_candidates("/random/path"), [])
def test_xwechat_files_at_end_returns_empty(self):
self.assertEqual(fkm.collect_wxid_candidates("/x/xwechat_files"), [])
class VerifyAesKeyTests(unittest.TestCase):
KEY = "a0c093edddc98490"
def _encrypt(self, plaintext_16):
return AES.new(self.KEY.encode("ascii"), AES.MODE_ECB).encrypt(plaintext_16)
def test_jpeg_magic_passes(self):
ct = self._encrypt(b"\xff\xd8\xff\xe0" + b"\x00" * 12)
self.assertTrue(fkm.verify_aes_key(self.KEY, ct))
def test_png_magic_passes(self):
ct = self._encrypt(b"\x89PNG\r\n\x1a\n" + b"\x00" * 8)
self.assertTrue(fkm.verify_aes_key(self.KEY, ct))
def test_gif_magic_passes(self):
ct = self._encrypt(b"GIF89a" + b"\x00" * 10)
self.assertTrue(fkm.verify_aes_key(self.KEY, ct))
def test_wxgf_magic_passes(self):
ct = self._encrypt(b"wxgf" + b"\x00" * 12)
self.assertTrue(fkm.verify_aes_key(self.KEY, ct))
def test_random_data_fails(self):
self.assertFalse(fkm.verify_aes_key(self.KEY, bytes(range(16))))
def test_wrong_length_template_fails(self):
self.assertFalse(fkm.verify_aes_key(self.KEY, b"short"))
self.assertFalse(fkm.verify_aes_key(self.KEY, b""))
def test_short_aes_key_fails(self):
self.assertFalse(fkm.verify_aes_key("short", b"\x00" * 16))
def test_empty_aes_key_fails(self):
self.assertFalse(fkm.verify_aes_key("", b"\x00" * 16))
class VerifyAesKeyAgainstAllTests(unittest.TestCase):
"""交叉验证:必须所有模板都通过才算命中(防短 magic 偶然碰撞)。"""
KEY = "a0c093edddc98490"
def _encrypt(self, plaintext_16):
return AES.new(self.KEY.encode("ascii"), AES.MODE_ECB).encrypt(plaintext_16)
def test_all_templates_pass(self):
ct1 = self._encrypt(b"\xff\xd8\xff\xe0" + b"\x00" * 12)
ct2 = self._encrypt(b"\x89PNG\r\n\x1a\n" + b"\x00" * 8)
self.assertTrue(fkm.verify_aes_key_against_all(self.KEY, [ct1, ct2]))
def test_one_template_fails_overall_fails(self):
ct1 = self._encrypt(b"\xff\xd8\xff\xe0" + b"\x00" * 12) # passes
ct2 = bytes(range(16)) # random, fails
self.assertFalse(fkm.verify_aes_key_against_all(self.KEY, [ct1, ct2]))
def test_empty_template_list_returns_false(self):
# 没模板就不能验证;不视为通过(防"零样本=自动通过"陷阱)
self.assertFalse(fkm.verify_aes_key_against_all(self.KEY, []))
class FindV2TemplateCiphertextsTests(unittest.TestCase):
def setUp(self):
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
self.dir = self._tmp.name
def _build_v2_dat(self, name, ciphertext_16, subdir=""):
target_dir = os.path.join(self.dir, subdir) if subdir else self.dir
os.makedirs(target_dir, exist_ok=True)
path = os.path.join(target_dir, name)
with open(path, "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + ciphertext_16 + b"\x00\x00")
return path
def test_finds_one_template_in_v2_thumb(self):
ct = bytes(range(0xF, 0x1F))
self._build_v2_dat("abc_t.dat", ct)
result = fkm.find_v2_template_ciphertexts(self.dir)
self.assertEqual(result, [ct])
def test_finds_multiple_distinct_templates(self):
cts = [bytes([i] * 16) for i in (0x11, 0x22, 0x33)]
for i, ct in enumerate(cts):
self._build_v2_dat(f"chat{i}_t.dat", ct, subdir=f"chat{i}")
result = fkm.find_v2_template_ciphertexts(self.dir, max_templates=3)
self.assertEqual(set(result), set(cts))
def test_dedupes_identical_templates(self):
ct = b"\x42" * 16
self._build_v2_dat("a_t.dat", ct, subdir="a")
self._build_v2_dat("b_t.dat", ct, subdir="b")
result = fkm.find_v2_template_ciphertexts(self.dir)
self.assertEqual(result, [ct])
def test_falls_back_to_any_dat_if_no_thumb(self):
ct = b"\x33" * 16
self._build_v2_dat("only_full.dat", ct)
self.assertEqual(fkm.find_v2_template_ciphertexts(self.dir), [ct])
def test_skips_non_v2_files(self):
path = os.path.join(self.dir, "abc_t.dat")
with open(path, "wb") as f:
f.write(b"\x00" * 100)
self.assertEqual(fkm.find_v2_template_ciphertexts(self.dir), [])
def test_empty_dir_returns_empty(self):
self.assertEqual(fkm.find_v2_template_ciphertexts(self.dir), [])
def test_missing_dir_returns_empty(self):
self.assertEqual(fkm.find_v2_template_ciphertexts("/nonexistent"), [])
def test_walks_into_subdirs(self):
ct = b"\x44" * 16
self._build_v2_dat("x_t.dat", ct, subdir="sub/deeper")
self.assertEqual(fkm.find_v2_template_ciphertexts(self.dir), [ct])
def test_respects_max_templates(self):
cts = [bytes([i] * 16) for i in range(10)]
for i, ct in enumerate(cts):
self._build_v2_dat(f"x{i}_t.dat", ct, subdir=f"d{i}")
result = fkm.find_v2_template_ciphertexts(self.dir, max_templates=2)
self.assertEqual(len(result), 2)
class FindImageKeyMacosIntegrationTests(unittest.TestCase):
"""端到端集成:合成 kvcomm 文件 + 合成 V2 模板 → 期望派生出已知 key。"""
def _build_test_env(self, tmpdir, code, wxid_raw, num_templates=2):
"""构造测试环境,返回 (db_dir, expected_xor, expected_aes)。"""
wxid_norm = fkm.normalize_wxid(wxid_raw)
base = os.path.join(tmpdir, "Documents", "xwechat_files", wxid_raw)
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
kvcomm = os.path.join(tmpdir, "Documents", "app_data", "net", "kvcomm")
os.makedirs(kvcomm)
with open(os.path.join(kvcomm, f"key_{code}_x.statistic"), "w") as f:
f.write("")
xor_expected, aes_expected = fkm.derive_image_keys(code, wxid_norm)
# 多个模板用不同的 plaintext 加密(仍是图像 magic 开头但内容不同)
plaintexts = [
b"\xff\xd8\xff\xe0" + b"\x00" * 12, # JPEG
b"\x89PNG\r\n\x1a\n" + b"\x00" * 8, # PNG
b"GIF89a" + b"\x01\x02" + b"\x00" * 8, # GIF
]
for i in range(num_templates):
pt = plaintexts[i % len(plaintexts)]
ct = AES.new(aes_expected.encode("ascii"), AES.MODE_ECB).encrypt(pt)
attach = os.path.join(base, "msg", "attach", f"chat{i}")
os.makedirs(attach)
with open(os.path.join(attach, f"img{i}_t.dat"), "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + ct + b"\x00\x00")
return db_dir, xor_expected, aes_expected
def test_full_flow_succeeds_with_normalized_wxid(self):
with tempfile.TemporaryDirectory() as tmp:
db_dir, xor_exp, aes_exp = self._build_test_env(
tmp, code=12345678, wxid_raw="your_wxid_a1b2", num_templates=3)
result = fkm.find_image_key_macos(db_dir)
self.assertIsNotNone(result, "派生应该成功")
self.assertEqual(result, (xor_exp, aes_exp))
def test_returns_none_when_no_kvcomm_codes(self):
with tempfile.TemporaryDirectory() as tmp:
base = os.path.join(tmp, "Documents", "xwechat_files", "wxid_x")
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
self.assertIsNone(fkm.find_image_key_macos(db_dir))
def test_returns_none_when_no_v2_template(self):
with tempfile.TemporaryDirectory() as tmp:
base = os.path.join(tmp, "Documents", "xwechat_files", "wxid_x")
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
kvcomm = os.path.join(tmp, "Documents", "app_data", "net", "kvcomm")
os.makedirs(kvcomm)
with open(os.path.join(kvcomm, "key_42_x.statistic"), "w") as f:
f.write("")
self.assertIsNone(fkm.find_image_key_macos(db_dir))
def test_returns_none_when_no_combination_verifies(self):
# 有 code 也有 V2 .dat但密文是随机的没有任何 key 能解出
with tempfile.TemporaryDirectory() as tmp:
base = os.path.join(tmp, "Documents", "xwechat_files", "wxid_x")
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
kvcomm = os.path.join(tmp, "Documents", "app_data", "net", "kvcomm")
os.makedirs(kvcomm)
with open(os.path.join(kvcomm, "key_42_x.statistic"), "w") as f:
f.write("")
attach = os.path.join(base, "msg", "attach", "x")
os.makedirs(attach)
with open(os.path.join(attach, "x_t.dat"), "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + b"\xde\xad\xbe\xef" * 4 + b"\x00\x00")
self.assertIsNone(fkm.find_image_key_macos(db_dir))
def test_empty_db_dir_returns_none_without_crash(self):
# 防御:空字符串、不合理路径不应抛异常。
# patch expanduser 让 HOME fallback 也指向不存在的路径,避免
# 测试在装了真实微信的开发机上意外深入到 wxid 缺失分支。
with tempfile.TemporaryDirectory() as fake_home:
with patch("os.path.expanduser", return_value=fake_home):
self.assertIsNone(fkm.find_image_key_macos(""))
class MainShortCircuitTests(unittest.TestCase):
"""main() 短路:已有 image_aes_key 仍然有效时,不应重新派生 / 不应改写 config。"""
def test_existing_valid_key_skips_derivation(self):
with tempfile.TemporaryDirectory() as tmp:
wxid = "wxid_abc"
base = os.path.join(tmp, "Documents", "xwechat_files", wxid)
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
# kvcomm 里放个 code证明若真去派生也能算出 key
kvcomm = os.path.join(tmp, "Documents", "app_data", "net", "kvcomm")
os.makedirs(kvcomm)
code = 42
with open(os.path.join(kvcomm, f"key_{code}_x.statistic"), "w") as f:
f.write("")
# 用真实派生的 key 加密 V2 模板,使现有 key 在该模板上能验证通过
xor_exp, aes_exp = fkm.derive_image_keys(code, wxid)
jpeg_pt = b"\xff\xd8\xff\xe0" + b"\x00" * 12
ct = AES.new(aes_exp.encode("ascii"), AES.MODE_ECB).encrypt(jpeg_pt)
attach = os.path.join(base, "msg", "attach", "x")
os.makedirs(attach)
with open(os.path.join(attach, "test_t.dat"), "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + ct + b"\x00\x00")
# 写入"已有有效 key"的 config
cfg_path = os.path.join(tmp, "config.json")
cfg_initial = {
"db_dir": db_dir,
"image_aes_key": aes_exp,
"image_xor_key": xor_exp,
"extra_field": "must_be_preserved", # 证明 main 不会重写
}
with open(cfg_path, "w", encoding="utf-8") as f:
json.dump(cfg_initial, f)
mtime_before = os.path.getmtime(cfg_path)
# 关键patch find_image_key_macos 让它若被误调用立刻可见
with patch.object(fkm, "find_image_key_macos") as mock_derive:
fkm.main(config_path=cfg_path)
mock_derive.assert_not_called() # 短路应直接 return不进派生
# config.json 不应被重写
self.assertEqual(os.path.getmtime(cfg_path), mtime_before)
with open(cfg_path, encoding="utf-8") as f:
self.assertEqual(json.load(f), cfg_initial)
def test_existing_invalid_key_falls_through_to_derivation(self):
with tempfile.TemporaryDirectory() as tmp:
wxid = "wxid_abc"
base = os.path.join(tmp, "Documents", "xwechat_files", wxid)
db_dir = os.path.join(base, "db_storage")
os.makedirs(db_dir)
kvcomm = os.path.join(tmp, "Documents", "app_data", "net", "kvcomm")
os.makedirs(kvcomm)
code = 42
with open(os.path.join(kvcomm, f"key_{code}_x.statistic"), "w") as f:
f.write("")
xor_exp, aes_exp = fkm.derive_image_keys(code, wxid)
jpeg_pt = b"\xff\xd8\xff\xe0" + b"\x00" * 12
ct = AES.new(aes_exp.encode("ascii"), AES.MODE_ECB).encrypt(jpeg_pt)
attach = os.path.join(base, "msg", "attach", "x")
os.makedirs(attach)
with open(os.path.join(attach, "test_t.dat"), "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + ct + b"\x00\x00")
cfg_path = os.path.join(tmp, "config.json")
cfg_initial = {
"db_dir": db_dir,
"image_aes_key": "deadbeefdeadbeef", # 故意写一个错的
}
with open(cfg_path, "w", encoding="utf-8") as f:
json.dump(cfg_initial, f)
fkm.main(config_path=cfg_path)
# 短路应失败,进入派生路径,配置应被改写为正确的 key
with open(cfg_path, encoding="utf-8") as f:
cfg_after = json.load(f)
self.assertEqual(cfg_after["image_aes_key"], aes_exp)
self.assertEqual(cfg_after["image_xor_key"], xor_exp)
# ---------- 方案2 (wxid 后缀候选搜索, fallback) 单元测试 ---------- #
class ExtractWxidPartsTests(unittest.TestCase):
"""extract_wxid_parts: 从 db_dir 提 (full, norm, suffix)。"""
def test_extracts_norm_and_suffix_from_alnum_suffix(self):
db_dir = "/foo/Documents/xwechat_files/your_wxid_25d5/db_storage"
self.assertEqual(
fkm.extract_wxid_parts(db_dir),
("your_wxid_25d5", "your_wxid", "25d5"),
)
def test_wxid_format_with_4char_suffix(self):
db_dir = "/foo/Documents/xwechat_files/wxid_abc_e2f4/db_storage"
self.assertEqual(
fkm.extract_wxid_parts(db_dir),
("wxid_abc_e2f4", "wxid_abc", "e2f4"),
)
def test_uppercase_suffix_lowercased(self):
db_dir = "/foo/Documents/xwechat_files/your_wxid_ABCD/db_storage"
result = fkm.extract_wxid_parts(db_dir)
self.assertIsNotNone(result)
self.assertEqual(result[2], "abcd")
def test_no_4char_suffix_returns_none(self):
# 6字符尾缀不匹配 _<4字符>$, 算法假设破灭
db_dir = "/foo/Documents/xwechat_files/wxid_simple/db_storage"
self.assertIsNone(fkm.extract_wxid_parts(db_dir))
def test_no_xwechat_files_returns_none(self):
self.assertIsNone(fkm.extract_wxid_parts("/random/path/db_storage"))
class DeriveXorKeyFromV2DatTests(unittest.TestCase):
"""derive_xor_key_from_v2_dat: 末字节投票反推 xor_key。"""
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.dir = self.tmp.name
def tearDown(self):
self.tmp.cleanup()
def _write_v2_dat(self, name, last_byte, subdir=""):
d = os.path.join(self.dir, subdir) if subdir else self.dir
os.makedirs(d, exist_ok=True)
body = (fkm.V2_MAGIC + b"\x00" * 9 + b"\x11" * 16
+ b"\x00" * 4 + bytes([last_byte]))
with open(os.path.join(d, name), "wb") as f:
f.write(body)
def test_unanimous_vote(self):
# 全部末字节 = 0xA6 → xor_key = 0xA6 ^ 0xD9 = 0x7F
for i in range(10):
self._write_v2_dat(f"x{i}_t.dat", 0xA6)
self.assertEqual(fkm.derive_xor_key_from_v2_dat(self.dir),
(0x7F, 10, 10))
def test_majority_vote_with_dissent(self):
# 8 个 0xA6, 2 个 0x55: 多数 0x7F 胜出
for i in range(8):
self._write_v2_dat(f"good{i}_t.dat", 0xA6)
for i in range(2):
self._write_v2_dat(f"bad{i}_t.dat", 0x55)
result = fkm.derive_xor_key_from_v2_dat(self.dir)
self.assertEqual(result, (0x7F, 8, 10))
def test_no_v2_dat_returns_none(self):
self.assertIsNone(fkm.derive_xor_key_from_v2_dat(self.dir))
def test_below_min_samples_returns_none(self):
# 默认 min_samples=3, 仅有 2 个样本应被视为不可信
for i in range(2):
self._write_v2_dat(f"x{i}_t.dat", 0xA6)
self.assertIsNone(fkm.derive_xor_key_from_v2_dat(self.dir))
def test_missing_dir_returns_none(self):
self.assertIsNone(fkm.derive_xor_key_from_v2_dat("/nonexistent"))
def test_walks_into_subdirs(self):
for i in range(10):
self._write_v2_dat(f"x{i}_t.dat", 0xA6, subdir=f"deep/sub{i}")
result = fkm.derive_xor_key_from_v2_dat(self.dir)
self.assertIsNotNone(result)
self.assertEqual(result[0], 0x7F)
def test_skips_non_v2_files(self):
# 不是 V2 magic 的 .dat 不计入投票
with open(os.path.join(self.dir, "junk.dat"), "wb") as f:
f.write(b"NOT_V2" + b"\x00" * 30)
for i in range(10):
self._write_v2_dat(f"x{i}_t.dat", 0xA6)
result = fkm.derive_xor_key_from_v2_dat(self.dir)
self.assertEqual(result, (0x7F, 10, 10))
class BruteforceUinCandidatesTests(unittest.TestCase):
"""bruteforce_uin_candidates: 候选枚举 + md5 前缀匹配。
注意test_real_bruteforce_against_golden 单核 ~7-8 秒,全套测试耗时大头。
"""
def test_real_bruteforce_against_golden(self):
# 真跑全空间 2^24 候选, 同时验证: (a) 合成 uin 在结果里
# (b) 候选数合理 (~256) (c) 候选都满足 xor_key 约束
# md5("12345678")[:4] == "25d5", 12345678 & 0xff == 0x4E
out = fkm.bruteforce_uin_candidates(0x4E, "25d5")
self.assertIn(12345678, out, "合成 uin 应在候选里")
self.assertTrue(200 <= len(out) <= 350,
f"候选数 {len(out)} 偏离 ~256 (理论 2^24/2^16)")
for uin in out[:20]:
self.assertEqual(uin & 0xFF, 0x4E,
f"uin {uin} 不满足 xor_key 约束")
class FindViaBruteforceTests(unittest.TestCase):
"""方案2 端到端 (合成 fixture, 多进程 worker 实跑)。
注: parallel 路径在合成 uin (低数值, 在 worker 0 chunk 早期命中) 上
< 0.2s 完成, 不需要 mock 加速。worker spawn 开销是真实集成测试的合理代价。
"""
def _build_bruteforce_env(self, tmp, uin, wxid_norm, suffix):
wxid_full = f"{wxid_norm}_{suffix}"
base = os.path.join(tmp, "Documents", "xwechat_files", wxid_full)
db_dir = os.path.join(base, "db_storage")
attach_dir = os.path.join(base, "msg", "attach")
os.makedirs(db_dir)
os.makedirs(attach_dir)
xor_exp, aes_exp = fkm.derive_image_keys(uin, wxid_norm)
last_byte = 0xD9 ^ xor_exp
jpeg_pt = b"\xff\xd8\xff\xe0" + b"\x00" * 12
ct = AES.new(aes_exp.encode("ascii"), AES.MODE_ECB).encrypt(jpeg_pt)
# 构造 10 个 V2 dat 让 derive_xor_key 投票稳定
for i in range(10):
with open(os.path.join(attach_dir, f"img{i}_t.dat"), "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + ct
+ b"\x00" * 4 + bytes([last_byte]))
return db_dir, attach_dir, xor_exp, aes_exp
def test_full_flow_finds_synthetic_uin(self):
with tempfile.TemporaryDirectory() as tmp:
db_dir, attach_dir, xor_exp, aes_exp = self._build_bruteforce_env(
tmp, uin=12345678, wxid_norm="your_wxid", suffix="25d5")
templates = fkm.find_v2_template_ciphertexts(attach_dir)
result = fkm._find_via_bruteforce(db_dir, attach_dir, templates)
self.assertEqual(result, (xor_exp, aes_exp))
def test_returns_none_when_no_wxid_suffix(self):
with tempfile.TemporaryDirectory() as tmp:
base = os.path.join(tmp, "Documents", "xwechat_files", "wxid_nosuffix")
db_dir = os.path.join(base, "db_storage")
attach_dir = os.path.join(base, "msg", "attach")
os.makedirs(db_dir)
os.makedirs(attach_dir)
self.assertIsNone(fkm._find_via_bruteforce(db_dir, attach_dir, []))
def test_returns_none_when_no_v2_dat(self):
with tempfile.TemporaryDirectory() as tmp:
base = os.path.join(tmp, "Documents", "xwechat_files",
"your_wxid_25d5")
db_dir = os.path.join(base, "db_storage")
attach_dir = os.path.join(base, "msg", "attach")
os.makedirs(db_dir)
os.makedirs(attach_dir)
self.assertIsNone(fkm._find_via_bruteforce(db_dir, attach_dir, []))
class DispatcherFallbackTests(unittest.TestCase):
"""find_image_key_macos dispatcher: 方案1 失败 → fallback 方案2。"""
def test_kvcomm_missing_falls_back_to_bruteforce(self):
with tempfile.TemporaryDirectory() as tmp:
uin, wxid_norm, suffix = 12345678, "your_wxid", "25d5"
wxid_full = f"{wxid_norm}_{suffix}"
base = os.path.join(tmp, "Documents", "xwechat_files", wxid_full)
db_dir = os.path.join(base, "db_storage")
attach_dir = os.path.join(base, "msg", "attach")
os.makedirs(db_dir)
os.makedirs(attach_dir)
# 不创建 kvcomm → 方案1 失败
xor_exp, aes_exp = fkm.derive_image_keys(uin, wxid_norm)
last_byte = 0xD9 ^ xor_exp
jpeg_pt = b"\xff\xd8\xff\xe0" + b"\x00" * 12
ct = AES.new(aes_exp.encode("ascii"), AES.MODE_ECB).encrypt(jpeg_pt)
for i in range(10):
with open(os.path.join(attach_dir, f"img{i}_t.dat"), "wb") as f:
f.write(fkm.V2_MAGIC + b"\x00" * 9 + ct
+ b"\x00" * 4 + bytes([last_byte]))
# patch HOME 让兜底 kvcomm 路径也找不到, 强制走方案2
with patch("os.path.expanduser", return_value=tmp):
result = fkm.find_image_key_macos(db_dir)
self.assertEqual(result, (xor_exp, aes_exp))
class BruteforceParallelTests(unittest.TestCase):
"""方案2 多进程实现的两层覆盖:
- 算法核心 (_bruteforce_worker_chunk): 直接调用, 无 process spawn, 极快
- 集成 (_bruteforce_with_aes_parallel): workers=1 验证 spawn + pickle 链路
多进程 e2e 由 FindViaBruteforceTests / DispatcherFallbackTests 间接覆盖
(cpu_count workers, 真实 fixture)。这里只测函数契约, 避免 spawn 开销
被反复支付。
"""
@classmethod
def setUpClass(cls):
# 合成 fixture, 跨多个测试复用
cls.uin = 12345678
cls.xor_key = cls.uin & 0xFF # 0x4E
cls.wxid_norm = "your_wxid"
cls.suffix_hex = hashlib.md5(str(cls.uin).encode()).hexdigest()[:4]
cls.suffix_bytes = bytes.fromhex(cls.suffix_hex)
cls.aes_hex = hashlib.md5(
f"{cls.uin}{cls.wxid_norm}".encode()
).hexdigest()[:16]
cls.template = AES.new(
cls.aes_hex.encode("ascii"), AES.MODE_ECB
).encrypt(b"\xff\xd8\xff\xe0" + b"\x00" * 12)
# i = (uin - xor_key) >> 8: worker 用 i 索引, 主进程倒推区间
cls.target_i = (cls.uin - cls.xor_key) >> 8
# 注: multiprocessing.Queue.put() 通过 feeder thread 异步刷到 pipe,
# get_nowait() 读取会 race。所有 queue 读用 get(timeout=...):
# - 命中场景: timeout=2s 给 feeder 充足时间 (实际 ~ms 级)
# - 不命中场景: timeout=0.5s 既证空又不拖慢测试
def test_worker_finds_known_uin_in_chunk(self):
q = multiprocessing.Queue()
fkm._bruteforce_worker_chunk(
self.target_i - 50, self.target_i + 50,
self.xor_key, self.suffix_bytes,
self.wxid_norm.encode("ascii"),
[self.template], q,
)
result = q.get(timeout=2)
self.assertEqual(result, (self.uin, self.aes_hex))
def test_worker_no_match_returns_silently(self):
# 区间不含 target_i (~48k), worker 扫完, queue 应保持空
q = multiprocessing.Queue()
fkm._bruteforce_worker_chunk(
0, 1000,
self.xor_key, self.suffix_bytes,
self.wxid_norm.encode("ascii"),
[self.template], q,
)
with self.assertRaises(_queue_mod.Empty):
q.get(timeout=0.5)
def test_worker_skips_when_aes_fails(self):
# md5 prefix 命中但 AES 模板错: 不入队 (防止 md5 单 gate 假阳)
q = multiprocessing.Queue()
wrong_template = b"\x00" * 16 # AES 解出来非图像 magic
fkm._bruteforce_worker_chunk(
self.target_i - 50, self.target_i + 50,
self.xor_key, self.suffix_bytes,
self.wxid_norm.encode("ascii"),
[wrong_template], q,
)
with self.assertRaises(_queue_mod.Empty):
q.get(timeout=0.5)
def test_parallel_workers_1_finds_synthetic_uin(self):
# 集成: workers=1 验证 process spawn + pickle + queue 跨进程通信
result = fkm._bruteforce_with_aes_parallel(
self.xor_key, self.suffix_hex, self.wxid_norm,
[self.template], workers=1, timeout=30,
)
self.assertEqual(result, (self.uin, self.aes_hex))
class SaveConfigAtomicTests(unittest.TestCase):
"""原子写测试os.replace 保证 config.json 不会被半截覆盖。"""
def test_roundtrip_writes_pretty_utf8(self):
with tempfile.TemporaryDirectory() as tmp:
cfg_path = os.path.join(tmp, "config.json")
cfg = {"db_dir": "/x", "image_aes_key": "中文测试key"}
fkm._save_config_atomic(cfg_path, cfg)
with open(cfg_path, encoding="utf-8") as f:
self.assertEqual(json.load(f), cfg)
# ensure_ascii=False中文应直接落盘不被转义
with open(cfg_path, "rb") as f:
self.assertIn("中文测试key".encode("utf-8"), f.read())
def test_failed_replace_leaves_original_intact(self):
with tempfile.TemporaryDirectory() as tmp:
cfg_path = os.path.join(tmp, "config.json")
with open(cfg_path, "w", encoding="utf-8") as f:
json.dump({"original": True}, f)
with patch.object(os, "replace",
side_effect=OSError("disk full during rename")):
with self.assertRaises(OSError):
fkm._save_config_atomic(cfg_path, {"new": True})
# 原文件应保持不变
with open(cfg_path, encoding="utf-8") as f:
self.assertEqual(json.load(f), {"original": True})
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,133 @@
"""Tests for `get_chat_images` multi-shard scanning.
WeChat rolls a chat's messages over to the next `message_N.db` shard once
the current one fills up, so any chat older than the current shard window
has its history split across multiple shards. The other query tools
(`get_chat_history`, `search_messages`, `decode_image`) already scan all
shards via `_find_msg_tables_for_user`; before this fix `get_chat_images`
used the single-shard `_find_msg_table_for_user`, so it silently dropped
every image that lived in a non-first shard.
These tests pin the corrected behaviour: results come from all matching
shards, are sorted by `create_time` DESC across shards, and respect the
`limit` cap.
"""
import unittest
from unittest.mock import patch
import mcp_server
class GetChatImagesMultiShardTests(unittest.TestCase):
def setUp(self):
# `resolve_username` / `get_contact_names` would hit real DBs; stub them.
self._patches = [
patch.object(mcp_server, "resolve_username",
side_effect=lambda x: "wxid_demo"),
patch.object(mcp_server, "get_contact_names",
return_value={"wxid_demo": "Demo"}),
]
for p in self._patches:
p.start()
self.addCleanup(p.stop)
def _run(self, shards, shard_images_map, limit=20):
"""Helper: stub the two collaborators and call the tool."""
def fake_list(db_path, table_name, username, limit=20, start_ts=None, end_ts=None):
return shard_images_map.get(db_path, [])
with patch.object(mcp_server, "_find_msg_tables_for_user",
return_value=shards), \
patch.object(mcp_server._image_resolver, "list_chat_images",
side_effect=fake_list):
return mcp_server.get_chat_images("Demo", limit=limit)
def test_collects_images_from_every_shard(self):
shards = [
{"db_path": "/m/message_1.db", "table_name": "Msg_x",
"max_create_time": 1_800_000_000},
{"db_path": "/m/message_2.db", "table_name": "Msg_x",
"max_create_time": 1_700_000_000},
]
shard_images = {
"/m/message_1.db": [
{"local_id": 11, "create_time": 1_800_000_000, "md5": "a" * 32, "size": 1024},
],
"/m/message_2.db": [
{"local_id": 22, "create_time": 1_700_000_000, "md5": "b" * 32, "size": 2048},
],
}
out = self._run(shards, shard_images)
# Both shards' images must appear; before the fix the message_2.db
# image was silently dropped.
self.assertIn("local_id=11", out)
self.assertIn("local_id=22", out)
self.assertIn("2 张图片", out)
def test_global_sort_by_create_time_desc(self):
# Older shard happens to contain a NEWER image (e.g. when shards are
# ordered by max_create_time but individual rows interleave): the
# output must still be globally sorted, not per-shard concatenated.
shards = [
{"db_path": "/m/message_1.db", "table_name": "Msg_x",
"max_create_time": 1_800_000_000},
{"db_path": "/m/message_2.db", "table_name": "Msg_x",
"max_create_time": 1_700_000_000},
]
shard_images = {
"/m/message_1.db": [
{"local_id": 11, "create_time": 1_750_000_000, "md5": "a" * 32},
],
"/m/message_2.db": [
# Older shard, but this single image is newer than the one above.
{"local_id": 22, "create_time": 1_799_000_000, "md5": "b" * 32},
],
}
out = self._run(shards, shard_images)
pos_22 = out.find("local_id=22")
pos_11 = out.find("local_id=11")
self.assertGreaterEqual(pos_22, 0)
self.assertGreaterEqual(pos_11, 0)
self.assertLess(pos_22, pos_11) # newer first
def test_limit_truncates_globally_across_shards(self):
shards = [
{"db_path": "/m/message_1.db", "table_name": "Msg_x",
"max_create_time": 1_800_000_000},
{"db_path": "/m/message_2.db", "table_name": "Msg_x",
"max_create_time": 1_700_000_000},
]
shard_images = {
"/m/message_1.db": [
{"local_id": i, "create_time": 1_800_000_000 - i}
for i in range(0, 5)
],
"/m/message_2.db": [
{"local_id": 100 + i, "create_time": 1_700_000_000 - i}
for i in range(0, 5)
],
}
out = self._run(shards, shard_images, limit=3)
# 3 newest overall = local_id=0, 1, 2 (all from shard 1, but the
# decision is global, not "first shard wins").
self.assertIn("3 张图片", out)
self.assertIn("local_id=0", out)
self.assertIn("local_id=1", out)
self.assertIn("local_id=2", out)
self.assertNotIn("local_id=100", out)
def test_no_shards_returns_not_found(self):
out = self._run(shards=[], shard_images_map={})
self.assertIn("找不到", out)
def test_all_shards_empty_returns_no_images(self):
shards = [
{"db_path": "/m/message_1.db", "table_name": "Msg_x",
"max_create_time": 0},
]
out = self._run(shards, shard_images_map={"/m/message_1.db": []})
self.assertIn("无图片消息", out)
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,85 @@
"""测试 _resolve_msg_types 和 _build_message_filters 的 type_filter 路径。"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import mcp_server
def test_resolve_none_returns_no_filter():
assert mcp_server._resolve_msg_types(None) == (None, None)
def test_resolve_empty_returns_no_filter():
assert mcp_server._resolve_msg_types([]) == (None, None)
def test_resolve_single_text_type():
type_filter, err = mcp_server._resolve_msg_types(['text'])
assert err is None
assert type_filter == [1]
def test_resolve_multiple_types():
type_filter, err = mcp_server._resolve_msg_types(['image', 'voice', 'video'])
assert err is None
assert sorted(type_filter) == [3, 34, 43]
def test_file_alias_maps_to_app():
"""'file' 是常见叫法, 实际是 type=49 (app message)。"""
type_filter, err = mcp_server._resolve_msg_types(['file'])
assert err is None
assert type_filter == [49]
def test_case_insensitive_and_strip():
type_filter, err = mcp_server._resolve_msg_types([' Text ', 'IMAGE'])
assert err is None
assert sorted(type_filter) == [1, 3]
def test_unknown_type_returns_error():
type_filter, err = mcp_server._resolve_msg_types(['unknown'])
assert type_filter is None
assert err is not None
assert 'unknown' in err
assert 'text' in err # 错误提示列出可选值
def test_partial_unknown_aborts_whole():
"""混入一个未知类型时整体失败, 不偷偷过滤合法的。"""
type_filter, err = mcp_server._resolve_msg_types(['text', 'invalid_type'])
assert type_filter is None
assert 'invalid_type' in err
def test_build_filters_without_type_filter():
"""type_filter=None 时 SQL 不包含 local_type 子句。"""
clauses, params = mcp_server._build_message_filters()
assert not any('local_type' in c for c in clauses)
def test_build_filters_with_single_type():
clauses, params = mcp_server._build_message_filters(type_filter=[1])
assert any('local_type IN (?)' == c for c in clauses)
assert 1 in params
def test_build_filters_with_multiple_types():
clauses, params = mcp_server._build_message_filters(type_filter=[1, 3, 34])
type_clause = [c for c in clauses if 'local_type' in c][0]
assert type_clause == 'local_type IN (?,?,?)'
assert params == [1, 3, 34]
def test_build_filters_combines_with_time_and_keyword():
clauses, params = mcp_server._build_message_filters(
start_ts=1000, end_ts=2000, keyword='hello', type_filter=[1]
)
assert 'create_time >= ?' in clauses
assert 'create_time <= ?' in clauses
assert 'message_content LIKE ?' in clauses
assert any('local_type' in c for c in clauses)
assert params == [1000, 2000, '%hello%', 1]

View File

@@ -0,0 +1,73 @@
"""Tests for `_format_namecard_text` (msg_type=42 鉴定).
Before this helper, type=42 messages fell through the generic non-text branch
and emitted `[名片] <raw XML>`, dumping the full `<msg .../>` element including
antispamticket, biznamecardinfo and head-image URLs. Those tokens are PII that
should not be piped to downstream LLM / log systems.
These tests pin the new behaviour: a compact `[名片] <head>: <bio>` line,
without any source-only XML fields.
"""
import unittest
import mcp_server
# Realistic-shape sample with the noisy / sensitive attrs that used to leak.
_REAL_NAMECARD = (
'<msg username="wxid_friend_demo" nickname="李雷" '
'antispamticket="v2_abc123def456_should_not_leak" '
'fullpy="lilei" shortpy="LL" alias="" '
'imagestatus="3" scene="17" province="北京" city="海淀" sign="" '
'sex="1" certflag="0" certinfo="搬砖工人 / 业余摄影" '
'brandIconUrl="https://wx.qlogo.cn/should_not_leak" '
'bigheadimgurl="https://wx.qlogo.cn/should_not_leak_big" '
'smallheadimgurl="https://wx.qlogo.cn/should_not_leak_small" />'
)
class FormatNamecardTextTests(unittest.TestCase):
def test_compact_line_for_real_namecard(self):
out = mcp_server._format_namecard_text(_REAL_NAMECARD)
self.assertEqual(out, "[名片] 李雷: 搬砖工人 / 业余摄影")
def test_no_pii_or_url_in_output(self):
out = mcp_server._format_namecard_text(_REAL_NAMECARD)
self.assertNotIn("antispamticket", out)
self.assertNotIn("v2_abc123def456", out)
self.assertNotIn("qlogo.cn", out)
self.assertNotIn("brandIconUrl", out)
self.assertNotIn("headimgurl", out)
def test_official_account_marked(self):
xml = (
'<msg username="gh_some_official" nickname="Some Official Account" '
'certinfo="一个公众号" />'
)
out = mcp_server._format_namecard_text(xml)
self.assertEqual(
out, "[名片] Some Official Account (公众号 gh_some_official): 一个公众号"
)
def test_no_certinfo_falls_back_to_head_only(self):
xml = '<msg username="wxid_demo" nickname="韩梅梅" />'
out = mcp_server._format_namecard_text(xml)
self.assertEqual(out, "[名片] 韩梅梅")
def test_only_username_when_nickname_missing(self):
xml = '<msg username="wxid_demo" nickname="" />'
out = mcp_server._format_namecard_text(xml)
self.assertEqual(out, "[名片] wxid_demo")
def test_missing_both_identifiers_returns_none(self):
xml = '<msg nickname="" username="" />'
self.assertIsNone(mcp_server._format_namecard_text(xml))
def test_broken_xml_returns_none(self):
self.assertIsNone(mcp_server._format_namecard_text(""))
self.assertIsNone(mcp_server._format_namecard_text("<msg "))
self.assertIsNone(mcp_server._format_namecard_text("not xml at all"))
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,116 @@
"""
issue #59: opt-in OpenAI Whisper API 后端的两条关键回归测试。
只测两件事:
1. 隐私契约: 文件 > 25MB 在调用 OpenAI SDK 之前就被拒绝(保证不会无意上传)
2. 缓存正确性: backend 不匹配的旧条目不会被命中(避免切后端时返回错后端结果)
其余路径要么琐碎默认值读取、要么坏掉时声音很大SDK 错误、ImportError
不再单独覆盖。
"""
import os
import sys
import tempfile
import unittest
from unittest.mock import MagicMock, patch
import mcp_server
class _CacheIsolationMixin:
"""与 test_voice_transcription_cache.py 同款隔离:避免污染 module-level 缓存状态。"""
def setUp(self):
self._saved_cache = mcp_server._voice_transcription_cache
self._saved_path = mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE
self._saved_warned = mcp_server._voice_transcription_save_warned
mcp_server._voice_transcription_cache = None
mcp_server._voice_transcription_save_warned = False
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE = os.path.join(
self._tmp.name, "voice_transcriptions.json"
)
def tearDown(self):
mcp_server._voice_transcription_cache = self._saved_cache
mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE = self._saved_path
mcp_server._voice_transcription_save_warned = self._saved_warned
class OpenAIBackendPrivacyTests(unittest.TestCase):
"""隐私契约:超限文件必须在 OpenAI SDK 实例化之前就被拒绝。
若有人把 size check 移到 OpenAI(api_key=...) 之后(即便仍在 upload 前),
本测试会失败 —— 这层防御边界值得守住。
"""
def test_oversize_audio_rejected_before_sdk_call(self):
# 写一个 26MB 临时 WAV (用稀疏写法快速生成)
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
f.seek(26 * 1024 * 1024)
f.write(b"\0")
big_path = f.name
self.addCleanup(os.unlink, big_path)
# 注入一个假的 openai 模块,保证 import 成功OpenAI 构造函数若被调用即测试失败
fake_openai = MagicMock()
fake_openai.OpenAI = MagicMock(
side_effect=AssertionError("OpenAI() must not be instantiated for oversize files")
)
fake_openai.AuthenticationError = type("AuthenticationError", (Exception,), {})
fake_openai.RateLimitError = type("RateLimitError", (Exception,), {})
fake_openai.APIError = type("APIError", (Exception,), {})
with patch.dict(sys.modules, {"openai": fake_openai}):
with self.assertRaises(RuntimeError) as ctx:
mcp_server._transcribe_openai(big_path)
self.assertIn("25MB", str(ctx.exception))
fake_openai.OpenAI.assert_not_called()
class CacheBackendMatchTests(_CacheIsolationMixin, unittest.TestCase):
"""缓存正确性backend 不匹配 → 视为 miss避免切后端时返回错后端结果。"""
def test_cache_hit_requires_backend_match(self):
# 种入一条 openai 后端的缓存条目
key = mcp_server._voice_transcription_cache_key("wxid_test", 42)
cache = mcp_server._load_voice_transcription_cache()
cache[key] = {
"text": "openai-result",
"language": "zh",
"create_time": 1700000000,
"backend": "openai",
"model_size": "whisper-1",
}
mcp_server._save_voice_transcription_cache()
# 当前后端是 local应当 miss → 走转录流程而非返回 "openai-result"
with patch.object(mcp_server, "TRANSCRIPTION_BACKEND", "local"), \
patch.object(mcp_server, "OPENAI_API_KEY", ""), \
patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.object(mcp_server, "_fetch_voice_row",
return_value=(b"\x02fake-silk-blob", 1700000001)), \
patch.object(mcp_server, "_silk_to_wav",
return_value=("/tmp/fake.wav", 24000 * 2)), \
patch.object(mcp_server, "_transcribe_local",
return_value={"text": "local-result", "language": "zh"}), \
patch.dict(sys.modules, {"whisper": MagicMock(), "pysilk": MagicMock()}):
result = mcp_server.transcribe_voice("test_contact", 42)
# 没返回旧 openai 缓存,而是走了 local 转录流程
self.assertNotIn("openai-result", result)
self.assertIn("local-result", result)
# 落盘的新条目应记录当前后端
mcp_server._voice_transcription_cache = None
reloaded = mcp_server._load_voice_transcription_cache()
self.assertEqual(reloaded[key]["backend"], "local")
self.assertEqual(reloaded[key]["text"], "local-result")
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,36 @@
"""测试分页提示语 _pagination_hint() 的边界行为。"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import mcp_server
def test_no_hint_when_count_less_than_limit():
"""count < limit 表示已读完当前条件下全部结果,不提示。"""
assert mcp_server._pagination_hint(count=10, limit=50, offset=0) == ""
def test_hint_when_count_equals_limit():
"""count == limit 时无法判断是否还有更多,提示下一页 offset。"""
hint = mcp_server._pagination_hint(count=50, limit=50, offset=0)
assert "可能还有更多" in hint
assert "offset=50" in hint
def test_hint_advances_offset_by_limit():
"""连续翻页时 offset 累加。"""
hint = mcp_server._pagination_hint(count=20, limit=20, offset=100)
assert "offset=120" in hint
def test_no_hint_when_limit_zero():
"""limit=0 是非法分页 (上游有 _validate_pagination 兜底);防御性返回空。"""
assert mcp_server._pagination_hint(count=0, limit=0, offset=0) == ""
def test_no_hint_when_count_exceeds_limit():
"""理论上 count > limit 不该发生 (调用方已 limit), 但若发生仍要提示。"""
hint = mcp_server._pagination_hint(count=51, limit=50, offset=0)
assert "可能还有更多" in hint

View File

@@ -0,0 +1,470 @@
"""Helper-level regression tests for the recorditem / decoder additions.
Focused on locking in the bugs fixed across PR #65's many review rounds so
they don't regress. Covers helpers that are easy to call in isolation:
- `_safe_basename` path-traversal sanitize (round-4 high #1)
- `_md5_file_chunked` streaming hash + size cap (round-6 medium #3)
- `_parse_message_content` group prefix stripping for both `:\n` and
`:<?xml`/`:<msg` shapes (round-7 high #1)
- `_parse_app_message_outer` retry-with-wider-limit only fires for
`<type>19</type>` content (round-5 medium #3)
- `_format_record_message_text` end-to-end expansion of a >20KB outer
type-19 message (round-5 high #1, round-2 P2-1)
- `_format_record_dataitem` per-datatype rendering for the 14 known
types incl. text / file / image / 视频号 etc.
The two MCP-tool wrappers (decode_file_message / decode_record_item) lean
heavily on module globals (WECHAT_BASE_DIR, _cache, MSG_DB_KEYS) and the
real wechat cache layout. They are exercised by real-data smoke runs in
the PR description rather than mocked here — mocking the entire wechat
cache tree would dwarf the actual logic under test.
"""
import hashlib
import os
import tempfile
import unittest
import mcp_server
# -------- _safe_basename ----------------------------------------------------
class SafeBasenameTests(unittest.TestCase):
def test_normal_filename_passes(self):
self.assertEqual(mcp_server._safe_basename('normal.pdf'), 'normal.pdf')
self.assertEqual(
mcp_server._safe_basename('Lec 4- 零和.pdf'), 'Lec 4- 零和.pdf'
)
self.assertEqual(
mcp_server._safe_basename('file (1).pdf'), 'file (1).pdf'
)
def test_absolute_path_rejected(self):
self.assertEqual(mcp_server._safe_basename('/etc/passwd'), '')
def test_parent_dir_rejected(self):
# Strict reject — should not return the basename 'sensitive'.
self.assertEqual(mcp_server._safe_basename('../../sensitive'), '')
self.assertEqual(mcp_server._safe_basename('..'), '')
def test_path_separator_rejected(self):
self.assertEqual(mcp_server._safe_basename('subdir/x.pdf'), '')
self.assertEqual(mcp_server._safe_basename('a\\b\\c.pdf'), '')
def test_nul_rejected(self):
self.assertEqual(mcp_server._safe_basename('with\x00nul.pdf'), '')
def test_empty_or_dot_rejected(self):
self.assertEqual(mcp_server._safe_basename(''), '')
self.assertEqual(mcp_server._safe_basename('.'), '')
def test_inner_dots_pass(self):
# 'file...with..dots.pdf' has no separator → fine.
self.assertEqual(
mcp_server._safe_basename('file...with..dots.pdf'),
'file...with..dots.pdf',
)
# -------- _md5_file_chunked -------------------------------------------------
class Md5FileChunkedTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.NamedTemporaryFile(delete=False)
self.tmp.write(b'x' * 1000)
self.tmp.close()
self.addCleanup(lambda: os.unlink(self.tmp.name))
def test_happy_path_matches_hashlib(self):
md5, err = mcp_server._md5_file_chunked(self.tmp.name)
self.assertIsNone(err)
self.assertEqual(md5, hashlib.md5(b'x' * 1000).hexdigest())
def test_size_cap_rejects_oversized_file(self):
md5, err = mcp_server._md5_file_chunked(self.tmp.name, max_size=500)
self.assertIsNone(md5)
self.assertIn('超过 md5 校验上限', err)
def test_missing_file_returns_error(self):
md5, err = mcp_server._md5_file_chunked('/tmp/no/such/path/here_xxx')
self.assertIsNone(md5)
self.assertIsNotNone(err)
# -------- _parse_message_content --------------------------------------------
class ParseMessageContentTests(unittest.TestCase):
def test_legacy_newline_prefix_in_group(self):
sender, text = mcp_server._parse_message_content(
'wxid_abc:\n<msg>hi</msg>', 1, is_group=True
)
self.assertEqual(sender, 'wxid_abc')
self.assertEqual(text, '<msg>hi</msg>')
def test_xml_decl_inline_prefix_in_group(self):
# round-7 high #1: 'sender:<?xml...' without newline
sender, text = mcp_server._parse_message_content(
'wxid_abc:<?xml version="1.0"?><msg>x</msg>', 1, is_group=True
)
self.assertEqual(sender, 'wxid_abc')
self.assertTrue(text.startswith('<?xml'))
def test_msg_inline_prefix_in_group(self):
sender, text = mcp_server._parse_message_content(
'wxid_abc:<msg>x</msg>', 1, is_group=True
)
self.assertEqual(sender, 'wxid_abc')
self.assertEqual(text, '<msg>x</msg>')
def test_private_chat_does_not_strip(self):
sender, text = mcp_server._parse_message_content(
'wxid_abc:<msg>x</msg>', 1, is_group=False
)
self.assertEqual(sender, '')
self.assertEqual(text, 'wxid_abc:<msg>x</msg>')
def test_bytes_content_returns_marker(self):
sender, text = mcp_server._parse_message_content(b'\x00\x01', 1, is_group=False)
self.assertEqual(sender, '')
self.assertEqual(text, '(二进制内容)')
# -------- _parse_app_message_outer ------------------------------------------
class ParseAppMessageOuterTests(unittest.TestCase):
def test_small_xml_uses_default_path(self):
outer = '<msg><appmsg><type>5</type><title>x</title></appmsg></msg>'
root = mcp_server._parse_app_message_outer(outer)
self.assertIsNotNone(root)
def test_oversized_non_record_xml_short_circuits(self):
# round-5 medium #3: only <type>19</type> content should retry under
# the wider 500K cap. A 25KB non-type-19 message must NOT be parsed
# under the wider limit.
outer = '<msg><appmsg><type>5</type><title>' + 'X' * 25000 + '</title></appmsg></msg>'
root = mcp_server._parse_app_message_outer(outer)
self.assertIsNone(root)
def test_oversized_record_xml_retries(self):
# type=19 content > 20KB should succeed under the wider cap.
big_desc = 'A' * 25000
outer = (
'<msg><appmsg><type>19</type><title>x</title>'
f'<recorditem><![CDATA[<recordinfo><title>x</title>'
f'<datalist count="1"><dataitem datatype="1">'
f'<datadesc>{big_desc}</datadesc></dataitem></datalist>'
f'</recordinfo>]]></recorditem></appmsg></msg>'
)
self.assertGreater(len(outer), 20000)
root = mcp_server._parse_app_message_outer(outer)
self.assertIsNotNone(root)
# -------- _format_record_dataitem ------------------------------------------
class FormatRecordDataitemTests(unittest.TestCase):
def _item(self, xml):
import xml.etree.ElementTree as ET
return ET.fromstring(xml)
def test_text(self):
item = self._item(
'<dataitem datatype="1"><datadesc>hello world</datadesc></dataitem>'
)
self.assertEqual(mcp_server._format_record_dataitem(item), 'hello world')
def test_file_with_title(self):
item = self._item(
'<dataitem datatype="8"><datatitle>report.pdf</datatitle></dataitem>'
)
self.assertEqual(
mcp_server._format_record_dataitem(item), '[文件] report.pdf'
)
def test_image(self):
item = self._item('<dataitem datatype="2"></dataitem>')
self.assertEqual(mcp_server._format_record_dataitem(item), '[图片]')
def test_finder_feed(self):
# round-2 datatype 22 视频号
item = self._item(
'<dataitem datatype="22"><finderFeed><desc>video desc</desc></finderFeed></dataitem>'
)
self.assertEqual(
mcp_server._format_record_dataitem(item), '[视频号] video desc'
)
def test_music(self):
item = self._item(
'<dataitem datatype="29"><datatitle>song</datatitle><datadesc>artist</datadesc></dataitem>'
)
self.assertEqual(
mcp_server._format_record_dataitem(item), '[音乐] song - artist'
)
def test_unknown_datatype_falls_back_to_desc(self):
item = self._item(
'<dataitem datatype="99"><datadesc>fallback content</datadesc></dataitem>'
)
self.assertEqual(
mcp_server._format_record_dataitem(item), 'fallback content'
)
def test_unknown_datatype_with_no_desc_uses_label(self):
item = self._item('<dataitem datatype="999"></dataitem>')
self.assertEqual(
mcp_server._format_record_dataitem(item), '[未知类型 999]'
)
# -------- _format_record_message_text end-to-end ---------------------------
class FormatRecordMessageTextTests(unittest.TestCase):
def _outer_with_items(self, items_xml, title='Big card', is_chatroom=False):
chatroom = '<isChatRoom>1</isChatRoom>' if is_chatroom else ''
recordinfo = (
f'<recordinfo><title>{title}</title>{chatroom}'
f'<datalist count="{items_xml.count("<dataitem")}">{items_xml}</datalist>'
f'</recordinfo>'
)
return (
'<?xml version="1.0"?><msg><appmsg><title>x</title><type>19</type>'
f'<recorditem><![CDATA[{recordinfo}]]></recorditem>'
'</appmsg></msg>'
)
def test_large_outer_expands_via_app_message_path(self):
# round-2 P2-1 + round-5 high #1: 大 outer 端到端必须能展开
items_xml = ''.join(
f'<dataitem datatype="1"><sourcename>S{i}</sourcename>'
f'<sourcetime>2025-01-01 00:00</sourcetime>'
f'<datadesc>{"X" * 600}</datadesc></dataitem>'
for i in range(40)
)
outer = self._outer_with_items(items_xml)
self.assertGreater(len(outer), 20000)
out = mcp_server._format_app_message_text(
outer,
(19 << 32) | 49,
False,
'wxid_dummy',
'dummy',
{},
)
self.assertIsNotNone(out)
self.assertIn('[聊天记录]', out)
self.assertIn('共 40 条', out)
# 每行带 0-based index
self.assertIn('[0] ', out)
self.assertIn('[1] ', out)
def test_empty_datalist_marks_loading(self):
# 空 datalist 应展示"(待加载)"而非"共 0 条"
outer = (
'<?xml version="1.0"?><msg><appmsg><title>x</title><type>19</type>'
'<recorditem><![CDATA[<recordinfo><title>x</title>'
'<isChatRoom>0</isChatRoom></recordinfo>]]></recorditem>'
'</appmsg></msg>'
)
out = mcp_server._format_app_message_text(
outer, (19 << 32) | 49, False, 'd', 'd', {}
)
self.assertIn('待加载', out)
def test_chatroom_marker_appended(self):
items_xml = (
'<dataitem datatype="1"><sourcename>A</sourcename>'
'<datadesc>hi</datadesc></dataitem>'
)
outer = self._outer_with_items(items_xml, title='G', is_chatroom=True)
out = mcp_server._format_app_message_text(
outer, (19 << 32) | 49, True, 'd', 'd', {}
)
self.assertIn('群聊转发', out)
def test_overflow_truncation_marker(self):
# > _RECORD_MAX_ITEMS dataitems should produce a
# "…(还有 N 条未显示)" line.
original_max = mcp_server._RECORD_MAX_ITEMS
try:
mcp_server._RECORD_MAX_ITEMS = 3
items_xml = ''.join(
f'<dataitem datatype="1"><datadesc>m{i}</datadesc></dataitem>'
for i in range(7)
)
outer = self._outer_with_items(items_xml)
out = mcp_server._format_app_message_text(
outer, (19 << 32) | 49, False, 'd', 'd', {}
)
self.assertIn('还有 4 条未显示', out)
finally:
mcp_server._RECORD_MAX_ITEMS = original_max
# -------- WCPay transfer (appmsg type=2000) --------------------------------
#
# All fixtures use synthetic placeholder values — no real wxid / fee / id /
# memo. paysubtype semantics are community consensus from open-source wechat
# tooling; treat any "未识别" branch as forward-compatible degradation.
class TransferPaysubTypeLabelTests(unittest.TestCase):
def test_known_subtypes_present(self):
labels = mcp_server._TRANSFER_PAYSUBTYPE_LABEL
self.assertEqual(labels['1'], '发起转账')
self.assertEqual(labels['3'], '已收款')
self.assertEqual(labels['4'], '已退还')
# 5/7/8 are version-dependent variants — locked to current text so a
# silent rename in mcp_server.py would surface here.
self.assertEqual(labels['5'], '过期已退还')
self.assertEqual(labels['7'], '待领取')
self.assertEqual(labels['8'], '已领取')
def _transfer_appmsg(
paysubtype='1',
fee_desc='¥100.00',
pay_memo='',
payer='wxid_payer_synth',
receiver='wxid_recv_synth',
transferid='1' + '0' * 27,
transcationid='1' + '0' * 27,
begin_ts='1746528000',
invalid_ts='1746614400',
title='微信转账',
des='请收钱',
feedesc_tag='feedesc',
paymemo_tag='pay_memo',
):
"""Build a synthetic appmsg type=2000 root element. All values are
placeholder; tests never run against real wechat data."""
import xml.etree.ElementTree as ET
fee_node = f'<{feedesc_tag}>{fee_desc}</{feedesc_tag}>' if fee_desc else ''
memo_node = f'<{paymemo_tag}>{pay_memo}</{paymemo_tag}>' if pay_memo else ''
xml_text = (
f'<msg><appmsg><title>{title}</title><des>{des}</des>'
f'<type>2000</type>'
f'<wcpayinfo>'
f'<paysubtype>{paysubtype}</paysubtype>'
f'{fee_node}{memo_node}'
f'<transferid>{transferid}</transferid>'
f'<transcationid>{transcationid}</transcationid>'
f'<begintransfertime>{begin_ts}</begintransfertime>'
f'<invalidtime>{invalid_ts}</invalidtime>'
f'<payer_username>{payer}</payer_username>'
f'<receiver_username>{receiver}</receiver_username>'
f'</wcpayinfo></appmsg></msg>'
)
return ET.fromstring(xml_text), xml_text
class ExtractTransferInfoTests(unittest.TestCase):
def test_full_fields_round_trip(self):
root, _ = _transfer_appmsg(paysubtype='3', pay_memo='lunch split')
appmsg = root.find('.//appmsg')
info = mcp_server._extract_transfer_info(appmsg)
self.assertIsNotNone(info)
self.assertEqual(info['paysubtype'], '3')
self.assertEqual(info['paysubtype_label'], '已收款')
self.assertEqual(info['fee_desc'], '¥100.00')
self.assertEqual(info['pay_memo'], 'lunch split')
self.assertEqual(info['payer_username'], 'wxid_payer_synth')
self.assertEqual(info['receiver_username'], 'wxid_recv_synth')
self.assertEqual(info['begin_transfer_time'], '1746528000')
self.assertEqual(info['invalid_time'], '1746614400')
self.assertTrue(info['transfer_id'].startswith('1'))
self.assertTrue(info['transcation_id'].startswith('1'))
def test_missing_wcpayinfo_returns_none(self):
import xml.etree.ElementTree as ET
root = ET.fromstring(
'<msg><appmsg><title>x</title><type>2000</type></appmsg></msg>'
)
appmsg = root.find('.//appmsg')
self.assertIsNone(mcp_server._extract_transfer_info(appmsg))
def test_camelcase_feedesc_falls_back(self):
# 部分微信版本字段名为 feeDesc 而非 feedesc
root, _ = _transfer_appmsg(feedesc_tag='feeDesc')
appmsg = root.find('.//appmsg')
info = mcp_server._extract_transfer_info(appmsg)
self.assertEqual(info['fee_desc'], '¥100.00')
def test_camelcase_paymemo_falls_back(self):
# paymemo (无下划线) 也是已知变体
root, _ = _transfer_appmsg(pay_memo='note', paymemo_tag='paymemo')
appmsg = root.find('.//appmsg')
info = mcp_server._extract_transfer_info(appmsg)
self.assertEqual(info['pay_memo'], 'note')
def test_unknown_paysubtype_label_degraded(self):
root, _ = _transfer_appmsg(paysubtype='99')
appmsg = root.find('.//appmsg')
info = mcp_server._extract_transfer_info(appmsg)
self.assertEqual(info['paysubtype'], '99')
self.assertIn('99', info['paysubtype_label'])
def test_empty_paysubtype_label_empty(self):
root, _ = _transfer_appmsg(paysubtype='')
appmsg = root.find('.//appmsg')
info = mcp_server._extract_transfer_info(appmsg)
self.assertEqual(info['paysubtype_label'], '')
class FormatTransferMessageTextTests(unittest.TestCase):
def test_initiate_with_amount(self):
root, _ = _transfer_appmsg(paysubtype='1')
appmsg = root.find('.//appmsg')
out = mcp_server._format_transfer_message_text(appmsg, '微信转账')
self.assertIn('[转账·发起转账]', out)
self.assertIn('¥100.00', out)
def test_received_with_memo(self):
root, _ = _transfer_appmsg(paysubtype='3', pay_memo='lunch')
appmsg = root.find('.//appmsg')
out = mcp_server._format_transfer_message_text(appmsg, '微信转账')
self.assertIn('[转账·已收款]', out)
self.assertIn('备注: lunch', out)
def test_missing_wcpayinfo_falls_back_to_title(self):
import xml.etree.ElementTree as ET
root = ET.fromstring(
'<msg><appmsg><title>微信转账</title><type>2000</type></appmsg></msg>'
)
appmsg = root.find('.//appmsg')
out = mcp_server._format_transfer_message_text(appmsg, '微信转账')
self.assertEqual(out, '[转账] 微信转账')
def test_missing_fee_desc_safe(self):
# 没有金额时也要给一行能看的输出,不能崩
root, _ = _transfer_appmsg(paysubtype='4', fee_desc='')
appmsg = root.find('.//appmsg')
out = mcp_server._format_transfer_message_text(appmsg, '微信转账')
self.assertIn('[转账·已退还]', out)
class AppMessageDispatchTransferTests(unittest.TestCase):
"""type=2000 must route through _format_transfer_message_text via
_format_app_message_text (so get_chat_history / export_chat both pick it up)."""
def test_dispatch_calls_transfer_helper(self):
_, xml_text = _transfer_appmsg(paysubtype='3', pay_memo='dinner')
out = mcp_server._format_app_message_text(
xml_text, 49, False, 'wxid_dummy', 'dummy', {}
)
self.assertIsNotNone(out)
self.assertIn('[转账·已收款]', out)
self.assertIn('¥100.00', out)
self.assertIn('备注: dinner', out)
if __name__ == '__main__':
unittest.main()

219
tests/test_refer_message.py Normal file
View File

@@ -0,0 +1,219 @@
"""微信引用回复消息appmsg type=57解析鉴定测试。
旧逻辑直接把 refermsg/content 按 [:160] 截断当摘要,对 type=3 (图片) /
34 (语音) / 47 (动画表情) / 49 (嵌套卡片) 这些"二进制"被引用消息会渲染
成 cdnurl + aeskey + md5 一坨乱码 (issue #44 #45)。本组测试 pin 新行为:
按 refer_type 给 schema-aware 摘要cdnurl / aeskey / md5 / cdnthumb /
voiceurl / externurl 全部不再泄漏到聊天历史。
合成 fixturewxid_synth_a / wxid_synth_b / 12345@chatroom / Sender A/B /
svrid 1 + 0*18无真实 PII。
"""
import unittest
import xml.etree.ElementTree as ET
import mcp_server
# ---------- 合成 fixture ----------
def _appmsg(refermsg_xml='', title='我的回复'):
"""组装一个最小 type=57 appmsg 元素。"""
xml = (
f'<msg><appmsg><type>57</type><title>{title}</title>'
f'{refermsg_xml}</appmsg></msg>'
)
root = ET.fromstring(xml)
return root.find('.//appmsg')
def _refermsg(refer_type, content, fromusr='wxid_synth_a',
displayname='Sender A', svrid='1' + '0' * 18,
chatusr='', createtime='1700000000'):
return (
'<refermsg>'
f'<type>{refer_type}</type>'
f'<svrid>{svrid}</svrid>'
f'<fromusr>{fromusr}</fromusr>'
f'<chatusr>{chatusr}</chatusr>'
f'<displayname>{displayname}</displayname>'
f'<createtime>{createtime}</createtime>'
f'<content>{content}</content>'
'</refermsg>'
)
# ---------- 标签映射 ----------
class ReferInnerTypeLabelTests(unittest.TestCase):
def test_known_refer_inner_labels(self):
self.assertEqual(mcp_server._REFER_INNER_TYPE_LABEL['3'], '图片')
self.assertEqual(mcp_server._REFER_INNER_TYPE_LABEL['34'], '语音')
self.assertEqual(mcp_server._REFER_INNER_TYPE_LABEL['47'], '动画表情')
self.assertEqual(mcp_server._REFER_INNER_TYPE_LABEL['49'], '链接/卡片')
def test_known_inner_appmsg_labels(self):
self.assertEqual(mcp_server._INNER_APPMSG_TYPE_LABEL['5'], '链接')
self.assertEqual(mcp_server._INNER_APPMSG_TYPE_LABEL['6'], '文件')
self.assertEqual(mcp_server._INNER_APPMSG_TYPE_LABEL['19'], '聊天记录')
# ---------- _extract_refer_info ----------
class ExtractReferInfoTests(unittest.TestCase):
def test_full_fields_round_trip(self):
appmsg = _appmsg(_refermsg('1', '原文本'), title='回复正文')
info = mcp_server._extract_refer_info(appmsg)
self.assertEqual(info['reply_text'], '回复正文')
self.assertEqual(info['refer_type'], '1')
self.assertEqual(info['refer_fromusr'], 'wxid_synth_a')
self.assertEqual(info['refer_displayname'], 'Sender A')
self.assertEqual(info['refer_svrid'], '1' + '0' * 18)
self.assertEqual(info['refer_content'], '原文本')
def test_missing_refermsg_returns_none(self):
appmsg = _appmsg(refermsg_xml='', title='孤儿回复')
self.assertIsNone(mcp_server._extract_refer_info(appmsg))
# ---------- _summarize_refer_content ----------
class SummarizeReferContentTests(unittest.TestCase):
def test_text_returns_original(self):
self.assertEqual(mcp_server._summarize_refer_content('1', '你好'), '你好')
def test_text_truncates_to_max_len(self):
long = '' * 200
out = mcp_server._summarize_refer_content('1', long, max_len=160)
self.assertEqual(len(out), 161) # 160 + '…'
self.assertTrue(out.endswith(''))
def test_image_returns_label_not_xml(self):
v2_image_xml = (
'<msg><img cdnthumburl="http://cdn.example/leak_thumb" '
'aeskey="aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" '
'md5="bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb" '
'cdnurl="http://cdn.example/leak_main" /></msg>'
)
out = mcp_server._summarize_refer_content('3', v2_image_xml)
self.assertEqual(out, '[图片]')
# PII / 二进制元数据不能泄漏到摘要
for leak in ('cdnurl', 'aeskey', 'md5', 'cdnthumb', 'leak_main'):
self.assertNotIn(leak, out)
def test_voice_returns_label(self):
v_xml = '<msg><voicemsg voicelength="3300" '\
'voiceurl="http://cdn.example/leak.silk" /></msg>'
out = mcp_server._summarize_refer_content('34', v_xml)
self.assertEqual(out, '[语音]')
self.assertNotIn('voiceurl', out)
def test_emoji_returns_label(self):
out = mcp_server._summarize_refer_content(
'47', '<msg><emoji md5="xx" externurl="leak.gif"/></msg>'
)
self.assertEqual(out, '[动画表情]')
self.assertNotIn('externurl', out)
self.assertNotIn('leak', out)
def test_nested_link_card_summary(self):
nested = '<msg><appmsg><type>5</type><title>分享标题</title>'\
'<url>http://example.com/leak</url></appmsg></msg>'
out = mcp_server._summarize_refer_content('49', nested)
self.assertEqual(out, '[链接] 分享标题')
self.assertNotIn('http', out)
self.assertNotIn('url', out)
def test_nested_record_card_summary(self):
nested = '<msg><appmsg><type>19</type><title>群聊天记录</title></appmsg></msg>'
out = mcp_server._summarize_refer_content('49', nested)
self.assertEqual(out, '[聊天记录] 群聊天记录')
def test_nested_invalid_xml_falls_back_to_card(self):
self.assertEqual(
mcp_server._summarize_refer_content('49', '<msg><appmsg'),
'[卡片]',
)
def test_unknown_refer_type_falls_back(self):
out = mcp_server._summarize_refer_content('999', 'irrelevant')
self.assertEqual(out, '[type=999]')
def test_empty_content_with_known_type(self):
self.assertEqual(mcp_server._summarize_refer_content('3', ''), '[图片]')
def test_xxe_payload_rejected_in_nested(self):
xxe = (
'<!DOCTYPE foo [<!ENTITY x SYSTEM "file:///etc/passwd">]>'
'<msg><appmsg><type>5</type><title>&x;</title></appmsg></msg>'
)
out = mcp_server._summarize_refer_content('49', xxe)
self.assertEqual(out, '[卡片]')
# ---------- _format_refer_message_text ----------
class FormatReferMessageTextTests(unittest.TestCase):
def _names(self):
return {'wxid_synth_a': 'Sender A', 'wxid_synth_b': 'Sender B'}
def test_text_refer_in_1v1(self):
appmsg = _appmsg(_refermsg('1', '你吃了吗'), title='吃了')
out = mcp_server._format_refer_message_text(
appmsg, is_group=False, chat_username='wxid_synth_a',
chat_display_name='Sender A', names=self._names(),
)
self.assertEqual(out, '吃了\n ↳ 回复 Sender A: 你吃了吗')
def test_image_refer_uses_label_not_xml_payload(self):
v2_image = (
'&lt;msg&gt;&lt;img cdnurl="leak" aeskey="leak" md5="leak"/&gt;&lt;/msg&gt;'
)
appmsg = _appmsg(_refermsg('3', v2_image), title='这张?')
out = mcp_server._format_refer_message_text(
appmsg, is_group=False, chat_username='wxid_synth_a',
chat_display_name='Sender A', names=self._names(),
)
self.assertIn('[图片]', out)
for leak in ('cdnurl', 'aeskey', 'md5'):
self.assertNotIn(leak, out)
def test_missing_refermsg_falls_back_to_title(self):
appmsg = _appmsg(refermsg_xml='', title='孤儿回复')
out = mcp_server._format_refer_message_text(
appmsg, is_group=False, chat_username='wxid_synth_a',
chat_display_name='Sender A', names={},
)
self.assertEqual(out, '孤儿回复')
def test_empty_reply_uses_placeholder(self):
appmsg = _appmsg(_refermsg('1', 'hi'), title='')
out = mcp_server._format_refer_message_text(
appmsg, is_group=False, chat_username='wxid_synth_a',
chat_display_name='Sender A', names=self._names(),
)
self.assertTrue(out.startswith('[引用消息]'))
# ---------- 调度入口 ----------
class AppMessageDispatchReferTests(unittest.TestCase):
def test_type57_dispatches_to_helper(self):
# _format_app_message_text 的 type=57 分支必须走 _format_refer_message_text
# 不再走旧的 inline [:160] 截断。
v2_image = '&lt;msg&gt;&lt;img cdnurl="leak_main"/&gt;&lt;/msg&gt;'
content = (
f'<msg><appmsg><type>57</type><title>看这个</title>'
f'{_refermsg("3", v2_image)}</appmsg></msg>'
)
out = mcp_server._format_app_message_text(
content, local_type=49, is_group=False,
chat_username='wxid_synth_a', chat_display_name='Sender A', names={},
)
self.assertIn('[图片]', out)
self.assertNotIn('leak_main', out)
self.assertNotIn('cdnurl', out)
if __name__ == '__main__':
unittest.main()

View File

@@ -0,0 +1,82 @@
"""Tests for `_format_voice_text` (msg_type=34 鉴定).
Voice messages previously rendered as a bare `[语音]` from the generic
non-text branch. LLMs reading chat history had no way to (a) judge whether a
clip was worth transcribing, or (b) call `decode_voice` without first round-
tripping through `get_voice_messages` to look up the `local_id`.
These tests pin the new behaviour: `[语音 Ns]` (duration to 1 decimal) when
the embedded `<voicemsg voicelength="">` is parseable, with graceful
fallback to `[语音]` on missing / zero / malformed length.
"""
import unittest
import mcp_server
def _voice_xml(length_ms):
return (
f'<msg><voicemsg endflag="1" length="2048" voicelength="{length_ms}" '
'clientmsgid="abc" fromusername="wxid_synth_a" '
'cancelflag="0" voiceformat="4" forwardflag="0" /></msg>'
)
class FormatVoiceTextTests(unittest.TestCase):
def test_renders_duration_with_one_decimal(self):
self.assertEqual(mcp_server._format_voice_text(_voice_xml(3300)), "[语音 3.3s]")
def test_subsecond_voice(self):
self.assertEqual(mcp_server._format_voice_text(_voice_xml(800)), "[语音 0.8s]")
def test_long_clip(self):
self.assertEqual(mcp_server._format_voice_text(_voice_xml(62000)), "[语音 62.0s]")
def test_missing_voicelength_falls_back(self):
xml = '<msg><voicemsg endflag="1" length="2048" /></msg>'
self.assertEqual(mcp_server._format_voice_text(xml), "[语音]")
def test_zero_voicelength_falls_back(self):
self.assertEqual(mcp_server._format_voice_text(_voice_xml(0)), "[语音]")
def test_non_numeric_voicelength_falls_back(self):
xml = '<msg><voicemsg voicelength="abc" /></msg>'
self.assertEqual(mcp_server._format_voice_text(xml), "[语音]")
def test_empty_content(self):
self.assertEqual(mcp_server._format_voice_text(""), "[语音]")
self.assertEqual(mcp_server._format_voice_text(None), "[语音]")
def test_missing_voicemsg_tag(self):
self.assertEqual(mcp_server._format_voice_text("<msg></msg>"), "[语音]")
def test_malformed_xml(self):
self.assertEqual(mcp_server._format_voice_text("<msg><voicemsg"), "[语音]")
def test_xxe_payload_rejected(self):
xxe = (
'<!DOCTYPE foo [<!ENTITY x SYSTEM "file:///etc/passwd">]>'
'<msg><voicemsg voicelength="1000" /></msg>'
)
self.assertEqual(mcp_server._format_voice_text(xxe), "[语音]")
def test_end_to_end_format_message_text_with_voicelength(self):
xml = _voice_xml(3300)
_, text = mcp_server._format_message_text(
local_id=72481, local_type=34, content=xml, is_group=False,
chat_username="wxid_synth_a", chat_display_name="A", names={},
create_time=1700000000,
)
self.assertEqual(text, "[语音 3.3s] (local_id=72481, ts=1700000000)")
def test_end_to_end_without_voicelength(self):
_, text = mcp_server._format_message_text(
local_id=99, local_type=34, content="<msg></msg>", is_group=False,
chat_username="wxid_synth_a", chat_display_name="A", names={},
create_time=0,
)
self.assertEqual(text, "[语音] (local_id=99)")
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,274 @@
import json
import os
import tempfile
import threading
import unittest
from unittest.mock import patch
import mcp_server
class _CacheIsolationMixin:
"""所有测试共享:隔离 module-level 缓存状态 + 指向 tempdir 的 cache 文件。"""
def setUp(self):
self._saved_cache = mcp_server._voice_transcription_cache
self._saved_path = mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE
self._saved_warned = mcp_server._voice_transcription_save_warned
mcp_server._voice_transcription_cache = None
mcp_server._voice_transcription_save_warned = False
self._tmp = tempfile.TemporaryDirectory()
self.addCleanup(self._tmp.cleanup)
mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE = os.path.join(
self._tmp.name, "voice_transcriptions.json"
)
def tearDown(self):
mcp_server._voice_transcription_cache = self._saved_cache
mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE = self._saved_path
mcp_server._voice_transcription_save_warned = self._saved_warned
class VoiceTranscriptionCachePersistenceTests(_CacheIsolationMixin, unittest.TestCase):
"""_load_voice_transcription_cache / _save_voice_transcription_cache 的持久化行为。"""
def test_load_missing_file_returns_empty_dict(self):
self.assertEqual(mcp_server._load_voice_transcription_cache(), {})
def test_save_and_reload_roundtrip(self):
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_foo:42"] = {
"text": "你好",
"language": "zh",
"create_time": 1700000000,
"model_size": "base",
}
mcp_server._save_voice_transcription_cache()
# 强制下一次 load 从磁盘读
mcp_server._voice_transcription_cache = None
reloaded = mcp_server._load_voice_transcription_cache()
self.assertEqual(reloaded["wxid_foo:42"]["text"], "你好")
self.assertEqual(reloaded["wxid_foo:42"]["language"], "zh")
def test_corrupt_file_returns_empty_dict(self):
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, "w", encoding="utf-8") as f:
f.write("{{ not valid json")
self.assertEqual(mcp_server._load_voice_transcription_cache(), {})
def test_non_dict_payload_returns_empty_dict(self):
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, "w", encoding="utf-8") as f:
json.dump(["not", "a", "dict"], f)
self.assertEqual(mcp_server._load_voice_transcription_cache(), {})
def test_utf8_preserved_on_disk(self):
# ensure_ascii=False 必须生效,否则中文会被转义成 \uXXXX
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_bar:1"] = {"text": "中文测试", "language": "zh"}
mcp_server._save_voice_transcription_cache()
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, "rb") as f:
raw = f.read()
self.assertIn("中文测试".encode("utf-8"), raw)
def test_save_without_prior_load_persists_empty_dict(self):
# 从未 load 过就直接 save应落盘一个空 dict而不是静默丢弃。
mcp_server._voice_transcription_cache = None
mcp_server._save_voice_transcription_cache()
self.assertTrue(os.path.exists(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE))
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, encoding="utf-8") as f:
self.assertEqual(json.load(f), {})
class VoiceTranscriptionCacheAtomicityTests(_CacheIsolationMixin, unittest.TestCase):
"""原子写 + crash-during-save 行为。"""
def test_write_is_atomic_via_rename(self):
# 先写入一份已有缓存
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_x:1"] = {"text": "initial", "language": "zh", "model_size": "base"}
mcp_server._save_voice_transcription_cache()
# 模拟:写 .tmp 正常但 os.replace 阶段失败
original_replace = os.replace
def flaky_replace(src, dst):
raise OSError("disk full during rename")
cache["wxid_x:1"] = {"text": "MUTATED", "language": "zh", "model_size": "base"}
with patch.object(os, "replace", side_effect=flaky_replace):
mcp_server._save_voice_transcription_cache() # 不应抛
# 磁盘上应仍然是 initial不是 MUTATED也不是损坏的半截文件
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, encoding="utf-8") as f:
disk = json.load(f)
self.assertEqual(disk["wxid_x:1"]["text"], "initial")
# .tmp 应该被清理,避免污染目录
tmp_path = mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE + ".tmp"
# 注patch 生效期间 os.replace 失败finally 里会尝试 unlink
_ = original_replace # 防 lint 警告
self.assertFalse(os.path.exists(tmp_path))
def test_early_save_error_preserves_existing_file(self):
# json.dump 在 .tmp 上抛异常时(模拟磁盘满 / 权限问题),主文件应保持原样;
# 注意此测试不是"写到一半中断"而是"写前就失败"的场景。
cache = mcp_server._load_voice_transcription_cache()
cache["wxid_y:1"] = {"text": "survives", "language": "zh", "model_size": "base"}
mcp_server._save_voice_transcription_cache()
cache["wxid_y:1"] = {"text": "DO NOT SEE", "language": "zh", "model_size": "base"}
def boom(*args, **kwargs):
raise OSError("disk full")
with patch.object(mcp_server.json, "dump", side_effect=boom):
mcp_server._save_voice_transcription_cache() # 静默降级,不抛
# 主文件没被破坏:仍然可 json.load 出原先内容
mcp_server._voice_transcription_cache = None
reloaded = mcp_server._load_voice_transcription_cache()
self.assertEqual(reloaded["wxid_y:1"]["text"], "survives")
class VoiceTranscriptionCacheConcurrencyTests(_CacheIsolationMixin, unittest.TestCase):
"""多线程下的 load/save 行为。"""
def test_concurrent_load_returns_same_dict_instance(self):
# 16 个线程同时触发首次 load应当只实际化一份 dictlock 生效)
barrier = threading.Barrier(16)
results = []
results_lock = threading.Lock()
def worker():
barrier.wait()
d = mcp_server._load_voice_transcription_cache()
with results_lock:
results.append(id(d))
threads = [threading.Thread(target=worker) for _ in range(16)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(len(set(results)), 1, "并发 load 应返回同一个 dict 对象")
def test_concurrent_save_does_not_corrupt(self):
# 多个线程同时 save磁盘上最终文件必须是合法 JSON原子写 + lock 保障)
cache = mcp_server._load_voice_transcription_cache()
for i in range(100):
cache[f"wxid_z:{i}"] = {
"text": f"msg-{i}",
"language": "zh",
"model_size": "base",
}
def worker():
mcp_server._save_voice_transcription_cache()
threads = [threading.Thread(target=worker) for _ in range(8)]
for t in threads:
t.start()
for t in threads:
t.join()
with open(mcp_server.VOICE_TRANSCRIPTION_CACHE_FILE, encoding="utf-8") as f:
disk = json.load(f) # 必须能解析
self.assertEqual(len(disk), 100)
class TranscribeVoiceCacheHitTests(_CacheIsolationMixin, unittest.TestCase):
"""transcribe_voice 的缓存命中 / 失效路径。"""
def _seed(self, key, entry):
cache = mcp_server._load_voice_transcription_cache()
cache[key] = entry
mcp_server._save_voice_transcription_cache()
def test_cache_hit_skips_fetch_and_transcribe(self):
key = mcp_server._voice_transcription_cache_key("wxid_test", 7)
self._seed(key, {
"text": "缓存命中文本",
"language": "zh",
"create_time": 1700000000,
"model_size": mcp_server.LOCAL_WHISPER_MODEL,
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test") as mock_resolve, \
patch.object(mcp_server, "_fetch_voice_row") as mock_fetch, \
patch.object(mcp_server, "_silk_to_wav") as mock_silk, \
patch.object(mcp_server, "_get_whisper_model") as mock_model:
result = mcp_server.transcribe_voice("test_contact", 7)
mock_resolve.assert_called_once_with("test_contact")
mock_fetch.assert_not_called()
mock_silk.assert_not_called()
mock_model.assert_not_called()
self.assertIn("缓存命中文本", result)
self.assertIn("(zh)", result)
def test_cache_hit_uses_placeholder_when_create_time_missing(self):
# 旧条目若没有 create_time 字段,不应崩溃
key = mcp_server._voice_transcription_cache_key("wxid_test", 8)
self._seed(key, {
"text": "历史条目",
"language": "zh",
"model_size": mcp_server.LOCAL_WHISPER_MODEL,
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.object(mcp_server, "_fetch_voice_row") as mock_fetch:
result = mcp_server.transcribe_voice("test_contact", 8)
mock_fetch.assert_not_called()
self.assertIn("历史条目", result)
def test_cache_hit_returns_empty_text_without_retranscribing(self):
# Whisper 返回空也要缓存;再次调用应直接返回空,不进入 miss 路径
key = mcp_server._voice_transcription_cache_key("wxid_test", 9)
self._seed(key, {
"text": "",
"language": "zh",
"create_time": 1700000000,
"model_size": mcp_server.LOCAL_WHISPER_MODEL,
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.object(mcp_server, "_fetch_voice_row") as mock_fetch:
result = mcp_server.transcribe_voice("test_contact", 9)
mock_fetch.assert_not_called()
self.assertIn("(zh)", result)
def test_model_mismatch_is_treated_as_miss(self):
# 缓存条目的 model_size 和当前 LOCAL_WHISPER_MODEL 不一致时,
# 不应命中;进入 miss 路径(这里无 whisper 依赖,应落到"缺少依赖"分支)。
key = mcp_server._voice_transcription_cache_key("wxid_test", 10)
self._seed(key, {
"text": "旧模型结果",
"language": "zh",
"create_time": 1700000000,
"model_size": "OUTDATED_MODEL",
})
with patch.object(mcp_server, "resolve_username", return_value="wxid_test"), \
patch.dict("sys.modules", {"whisper": None}):
# whisper=None 时 `import whisper` 触发 ImportError
result = mcp_server.transcribe_voice("test_contact", 10)
# 走了 miss 路径 → 返回缺依赖提示,而不是返回旧缓存文本
self.assertNotIn("旧模型结果", result)
self.assertIn("缺少依赖", result)
def test_cache_key_handles_colon_in_username(self):
# 若上游未来的 resolve_username 放出带 ':' 的 username也不会和其他条目冲突
key_a = mcp_server._voice_transcription_cache_key("wxid:foo", 1)
key_b = mcp_server._voice_transcription_cache_key("wxid", 1)
self.assertNotEqual(key_a, key_b)
if __name__ == "__main__":
unittest.main()