Files
kefu/wechat_rpa/test_archive_auto_backup.py
T
2026-08-27 14:04:28 +08:00

439 lines
16 KiB
Python

# -*- coding: utf-8 -*-
"""桌面启动自动归档桥接器的增量与素材关联测试。"""
import json
import sqlite3
import tempfile
from pathlib import Path
from unittest import TestCase, main, mock
from archive_content_parser import (
decode_hex_protobuf_text,
parse_file_message_metadata,
parse_mini_program_metadata,
)
from archive_auto_backup import (
ArchiveApiClient,
AutoBackupConfig,
IncrementalArchiveImporter,
_load_exporter_keys,
)
class _FakeExporter:
@staticmethod
def connect_sqlite(path):
connection = sqlite3.connect(path)
connection.text_factory = lambda value: value.decode("utf-8", errors="replace")
return connection
@staticmethod
def parse_content(value):
if isinstance(value, bytes):
return value.decode("utf-8", errors="replace")
return str(value or "")
@staticmethod
def get_msg_type_name(value):
return {2: "文本", 3: "图片", 20: "合并转发", 78: "修改群名"}.get(
int(value), "未知"
)
class _FakeMediaExporter:
def __init__(self, media_path: Path):
self.media_path = media_path
@staticmethod
def extract_media_refs(content):
if isinstance(content, bytes) and b"media-ref" in content:
return {"uuids": ["media-ref"], "urls": [], "filenames": []}
return {"uuids": [], "urls": [], "filenames": []}
@staticmethod
def build_cache_index(_account_dir):
return {"by_uuid": {}, "by_name": {}}
def match_media(self, refs, _index):
if refs["uuids"]:
return str(self.media_path), "Image", "uuid"
return None, None, None
class _FakeApi:
def __init__(self):
self.cursor = {}
self.payloads = []
self.uploads = []
self.metadata_payloads = []
def checkpoint(self, _account):
return dict(self.cursor)
def advance_checkpoint(
self, _account, checkpoint, *, display_name="", corp_scope_id=""
):
self.cursor = dict(checkpoint)
return dict(checkpoint)
def pending_attachments(self, _account, limit=500):
return []
def upload_media(self, path):
self.uploads.append(Path(path))
return "media-001"
def sync_metadata(self, payload):
self.metadata_payloads.append(payload)
return {
"people_synced": len(payload.get("people") or []),
"conversations_updated": len(payload.get("conversations") or []),
}
def import_messages(self, payload):
self.payloads.append(payload)
self.cursor = dict(payload["checkpoint"])
return {
"received": len(payload["messages"]),
"inserted": len(payload["messages"]),
"duplicates": 0,
}
class _FailingMediaApi(_FakeApi):
def upload_media(self, path):
raise RuntimeError(f"COS unavailable: {path}")
class ArchiveAutoBackupTest(TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.account = "16880001"
self.source = self.root / "WXWork"
(self.source / self.account).mkdir(parents=True)
self.decrypted = self.root / "decrypted" / self.account
self.decrypted.mkdir(parents=True)
self.message_db = self.decrypted / "message.db"
database = sqlite3.connect(self.message_db)
try:
database.execute(
"""CREATE TABLE message_table(
send_time INTEGER,conversation_id TEXT,sender_id TEXT,
content BLOB,content_type INTEGER,server_id TEXT,
client_id TEXT,sequence INTEGER)"""
)
database.executemany(
"INSERT INTO message_table VALUES (?,?,?,?,?,?,?,?)",
[
(1000, "S:16880001_20001", "20001", b"hello", 2,
"server-1", "client-1", 1),
(1001, "R:room-1", self.account, b"media-ref", 3,
"server-2", "client-2", 2),
],
)
database.commit()
finally:
database.close()
self.media = self.root / "source-image.png"
self.media.write_bytes(b"png-data")
self.config = AutoBackupConfig(
exporter_root=self.root,
source_root=self.source,
work_root=self.root / "work",
api_urls=("http://127.0.0.1:8766",),
corp_scope_id="corp-one",
batch_size=1,
)
def tearDown(self):
self.temp.cleanup()
def test_incremental_batches_upload_media_before_advancing_checkpoint(self):
api = _FakeApi()
importer = IncrementalArchiveImporter(
self.config,
api,
_FakeExporter(),
_FakeMediaExporter(self.media),
)
decrypted = [(str(self.message_db), "message.db", self.account)]
first = importer.run(decrypted)
second = importer.run(decrypted)
self.assertEqual(first["inserted"], 2)
self.assertEqual(first["batches"], 2)
self.assertEqual(first["media"], 1)
self.assertEqual(second["inserted"], 0)
self.assertEqual(len(api.uploads), 1)
self.assertEqual(api.payloads[-1]["messages"][0]["media_ids"], ["media-001"])
self.assertEqual(api.payloads[-1]["messages"][0]["direction"], "outbound")
self.assertEqual(
api.payloads[0]["messages"][0]["sender"]["scope_id"], "corp-one"
)
self.assertEqual(api.cursor, {"send_time": 1001.0, "rowid": 2})
def test_metadata_resolves_direct_chat_peer_nickname(self):
user_db = self.decrypted / "user.db"
with sqlite3.connect(user_db) as database:
database.execute(
"CREATE TABLE user_table(id TEXT,name TEXT,real_name TEXT,account TEXT)"
)
database.executemany(
"INSERT INTO user_table VALUES (?,?,?,?)",
[
(self.account, "归档账号", "", ""),
("20001", "客户昵称", "", ""),
],
)
database.close()
session_db = self.decrypted / "session.db"
with sqlite3.connect(session_db) as database:
database.execute(
"""CREATE TABLE conversation_table(
id TEXT,name TEXT,roomname_remark TEXT,session_id TEXT)"""
)
database.execute(
"INSERT INTO conversation_table VALUES (?,?,?,?)",
("S:16880001_20001", "", "", ""),
)
database.close()
api = _FakeApi()
importer = IncrementalArchiveImporter(
self.config,
api,
_FakeExporter(),
_FakeMediaExporter(self.media),
)
importer.run(
[
(str(user_db), "user.db", self.account),
(str(session_db), "session.db", self.account),
(str(self.message_db), "message.db", self.account),
]
)
self.assertEqual(
api.payloads[0]["messages"][0]["conversation"]["name"], "客户昵称"
)
self.assertEqual(
api.metadata_payloads[0]["conversations"][0]["name"], "客户昵称"
)
def test_exporter_key_file_is_always_read_as_utf8(self):
key = "ab" * 16
(self.root / "wxwork_keys.json").write_text(
json.dumps({"备注": "中文", "keys": {self.account: key}}, ensure_ascii=False),
encoding="utf-8",
)
self.assertEqual(_load_exporter_keys(self.root), {self.account: key})
def test_nested_short_text_protobuf_keeps_single_character_and_emoji(self):
cases = {
"0a07080012030a0131": "1",
"0a09080012050a03e5a5bd": "好",
"0a0a080012060a04f09f918c": "👌",
"0a08080012040a023131": "11",
}
for encoded, expected in cases.items():
with self.subTest(encoded=encoded):
self.assertEqual(
decode_hex_protobuf_text(encoded, "文本"), expected
)
self.assertEqual(decode_hex_protobuf_text("1688857886854158", "安全通知"), "")
def test_file_card_is_structured_and_marked_for_cache_retry(self):
def varint(value):
output = bytearray()
while value >= 0x80:
output.append((value & 0x7F) | 0x80)
value >>= 7
output.append(value)
return bytes(output)
def bytes_field(number, value):
return varint((number << 3) | 2) + varint(len(value)) + value
file_size = 162_862_195
filename = "DoctorWorkstation-Setup-Windows-x64-1.0.0.exe"
payload = b"".join(
[
bytes_field(1, b"opaque-source-reference"),
bytes_field(2, filename.encode()),
varint(4 << 3) + varint(file_size),
bytes_field(10, b"9D4295ED638B290925A5A960D644934E"),
]
)
metadata = parse_file_message_metadata(payload, 20)
self.assertEqual(metadata["original_filename"], filename)
self.assertEqual(metadata["size_bytes"], file_size)
database = sqlite3.connect(self.message_db)
try:
database.execute(
"INSERT INTO message_table VALUES (?,?,?,?,?,?,?,?)",
(1002, "S:16880001_20001", self.account, payload, 20,
"server-file", "client-file", 3),
)
database.commit()
finally:
database.close()
api = _FakeApi()
importer = IncrementalArchiveImporter(
self.config, api, _FakeExporter(), _FakeMediaExporter(self.media)
)
result = importer.run([(str(self.message_db), "message.db", self.account)])
message = api.payloads[-1]["messages"][0]
self.assertEqual(result["inserted"], 3)
self.assertEqual(message["message_type"], "文件")
self.assertIn(filename, message["content"])
self.assertIn("源文件未缓存", message["content"])
self.assertEqual(
message["attachment_metadata"][0]["status"], "source_not_cached"
)
def test_application_conversation_is_skipped_and_checkpoint_advances(self):
database = sqlite3.connect(self.message_db)
try:
database.execute(
"INSERT INTO message_table VALUES (?,?,?,?,?,?,?,?)",
(1002, "Y:10011", self.account, b"approval notice", 2,
"server-app", "client-app", 3),
)
database.commit()
finally:
database.close()
api = _FakeApi()
importer = IncrementalArchiveImporter(
self.config, api, _FakeExporter(), _FakeMediaExporter(self.media)
)
result = importer.run([(str(self.message_db), "message.db", self.account)])
imported_conversations = {
message["conversation"]["external_id"]
for payload in api.payloads
for message in payload["messages"]
}
self.assertNotIn("Y:10011", imported_conversations)
self.assertEqual(result["application_messages_skipped"], 1)
self.assertEqual(result["inserted"], 2)
self.assertEqual(api.cursor, {"send_time": 1002.0, "rowid": 3})
def test_mini_program_card_is_decoded_from_protobuf(self):
def varint(value):
output = bytearray()
while value >= 0x80:
output.append((value & 0x7F) | 0x80)
value >>= 7
output.append(value)
return bytes(output)
def bytes_field(number, value):
return varint((number << 3) | 2) + varint(len(value)) + value
page_path = b"pages/order/monad/monad.html?id=11&doctor_id=117"
nested = b"".join(
[
bytes_field(1, b"gh_example@app"),
bytes_field(2, b"wx79b9a0bfbfe7cbcd"),
bytes_field(3, page_path),
bytes_field(6, b"https://example.test/cover.png"),
bytes_field(7, "点击进入诊室".encode()),
bytes_field(10, "甄养堂互联网医院".encode()),
]
)
raw = bytes_field(3, "点击进入诊室".encode()) + bytes_field(107, nested)
metadata = parse_mini_program_metadata(raw, 78)
self.assertEqual(metadata["title"], "点击进入诊室")
self.assertEqual(metadata["app_name"], "甄养堂互联网医院")
self.assertEqual(metadata["page_path"], page_path.decode())
database = sqlite3.connect(self.message_db)
try:
database.execute(
"INSERT INTO message_table VALUES (?,?,?,?,?,?,?,?)",
(1002, "S:16880001_20001", self.account, raw, 78,
"server-mini", "client-mini", 3),
)
database.commit()
finally:
database.close()
api = _FakeApi()
importer = IncrementalArchiveImporter(
self.config, api, _FakeExporter(), _FakeMediaExporter(self.media)
)
importer.run([(str(self.message_db), "message.db", self.account)])
message = api.payloads[-1]["messages"][0]
self.assertEqual(message["message_type"], "小程序")
self.assertIn("甄养堂互联网医院", message["content"])
self.assertIn(page_path.decode(), message["content"])
def test_media_upload_failure_does_not_advance_that_batch_checkpoint(self):
api = _FailingMediaApi()
importer = IncrementalArchiveImporter(
AutoBackupConfig(
exporter_root=self.config.exporter_root,
source_root=self.config.source_root,
work_root=self.config.work_root,
api_urls=self.config.api_urls,
corp_scope_id=self.config.corp_scope_id,
batch_size=2,
),
api,
_FakeExporter(),
_FakeMediaExporter(self.media),
)
decrypted = [(str(self.message_db), "message.db", self.account)]
with self.assertRaisesRegex(RuntimeError, "COS unavailable"):
importer.run(decrypted)
self.assertEqual(api.payloads, [])
self.assertEqual(api.cursor, {})
def test_client_uploads_large_media_parts_and_reports_etags(self):
large = self.root / "video.mp4"
large.write_bytes(b"0123456789abcdefghij")
client = ArchiveApiClient("http://127.0.0.1:8766", "test-key")
prepared = {
"media": {"id": "media-large"},
"reused": False,
"upload_mode": "multipart",
"multipart": {
"upload_id": "upload-1",
"part_size": 10,
"parts": [
{"part_number": 1, "size_bytes": 10, "upload_url": "https://cos/1"},
{"part_number": 2, "size_bytes": 10, "upload_url": "https://cos/2"},
],
"completed_parts": [
{"part_number": 1, "size_bytes": 10, "etag": "etag-1"}
],
},
}
client._json = mock.Mock(side_effect=[prepared, {"media": {"status": "ready"}}])
def uploaded(url, data, timeout):
response = mock.Mock()
response.headers = {"ETag": f'"etag-{url[-1]}"'}
response.raise_for_status.return_value = None
return response
try:
with mock.patch(
"archive_auto_backup.requests.put", side_effect=uploaded
) as put:
self.assertEqual(client.upload_media(large), "media-large")
self.assertEqual(put.call_count, 1)
finally:
client.close()
completed = client._json.call_args_list[-1]
self.assertTrue(completed.args[1].endswith("/multipart-complete"))
self.assertEqual(
completed.kwargs["payload"]["parts"],
[
{"part_number": 1, "etag": "etag-1"},
{"part_number": 2, "etag": "etag-2"},
],
)
if __name__ == "__main__":
main()