Add Windows GUI and WXWork export support
This commit is contained in:
104
tests/test_chat_export_helpers.py
Normal file
104
tests/test_chat_export_helpers.py
Normal 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()
|
||||
97
tests/test_chat_images_query_align.py
Normal file
97
tests/test_chat_images_query_align.py
Normal 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
|
||||
481
tests/test_decode_image_v2.py
Normal file
481
tests/test_decode_image_v2.py
Normal 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()
|
||||
295
tests/test_decode_images_batch.py
Normal file
295
tests/test_decode_images_batch.py
Normal 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()
|
||||
770
tests/test_find_image_key_macos.py
Normal file
770
tests/test_find_image_key_macos.py
Normal 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()
|
||||
133
tests/test_get_chat_images_multishard.py
Normal file
133
tests/test_get_chat_images_multishard.py
Normal 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()
|
||||
85
tests/test_msg_types_filter.py
Normal file
85
tests/test_msg_types_filter.py
Normal 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]
|
||||
73
tests/test_namecard_format.py
Normal file
73
tests/test_namecard_format.py
Normal 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()
|
||||
116
tests/test_openai_backend.py
Normal file
116
tests/test_openai_backend.py
Normal 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()
|
||||
36
tests/test_pagination_hint.py
Normal file
36
tests/test_pagination_hint.py
Normal 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
|
||||
470
tests/test_record_decoders.py
Normal file
470
tests/test_record_decoders.py
Normal 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
219
tests/test_refer_message.py
Normal 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 全部不再泄漏到聊天历史。
|
||||
|
||||
合成 fixture:wxid_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 = (
|
||||
'<msg><img cdnurl="leak" aeskey="leak" md5="leak"/></msg>'
|
||||
)
|
||||
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 = '<msg><img cdnurl="leak_main"/></msg>'
|
||||
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()
|
||||
82
tests/test_voice_format.py
Normal file
82
tests/test_voice_format.py
Normal 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()
|
||||
274
tests/test_voice_transcription_cache.py
Normal file
274
tests/test_voice_transcription_cache.py
Normal 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,应当只实际化一份 dict(lock 生效)
|
||||
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()
|
||||
Reference in New Issue
Block a user