# -*- coding: utf-8 -*- """Synthetic desktop normalization, upload privacy and resumable repair tests.""" import hashlib from contextlib import closing import json import sqlite3 import tempfile import unittest from pathlib import Path from unittest import mock import admin_backend import archive_auto_backup as archive from archive_store import ArchiveStore from archive_wxwork_adapter import WxworkArchiveAdapter from test_archive_auto_backup import _FakeExporter, _FakeMediaExporter, _FakeApi class Api(_FakeApi): base_url = "http://archive.test" archive_scope = "test-account" def import_messages(self, payload): self.payloads.append(payload) if "checkpoint" in payload: self.cursor = dict(payload["checkpoint"]) return {"received": len(payload["messages"]), "inserted": 0, "duplicates": len(payload["messages"]), "errors": 0} class PlaintextArchiveTest(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.root = Path(self.temp.name) self.config = archive.AutoBackupConfig(self.root, self.root / "source", self.root / "work", ("http://archive.test",), batch_size=1) self.api = Api() self.importer = archive.IncrementalArchiveImporter(self.config, self.api, _FakeExporter(), _FakeMediaExporter(self.root / "unused.png")) self.account = "10001" self.row = {"__archive_rowid": 1, "send_time": 1000, "conversation_id": "M:20002", "sender_id": "20002", "content_type": 2, "content": bytes.fromhex("0a07080012030a0131"), "extra_content": b"private-session-token", "server_id": "s1", "client_id": "c1"} def tearDown(self): self.temp.cleanup() def normalize(self, row=None, voice_texts=None): return self.importer._normalize(self.account, row or self.row, {}, {}, voice_texts) def test_raw_proto_is_decoded_without_legacy_exporter_parser(self): with mock.patch.object(_FakeExporter, "parse_content", side_effect=AssertionError("legacy parser must not run")): message = self.normalize() self.assertEqual(message["content"], "1") self.assertEqual(message["content_parse_status"], "decoded") self.assertEqual(message["source_message_id"], "s1") self.assertEqual(message["raw_fields"]["content_type"], 2) def test_only_body_digest_is_transmitted_even_for_printable_body(self): row = dict(self.row, content="你好", extra_content="private-key-material", opaque_blob=b"binary-private-field") message = self.normalize(row) self.assertEqual(message["content"], "你好") serialized = json.dumps(message) self.assertNotIn("private-key-material", serialized) self.assertNotIn("binary-private-field", serialized) self.assertNotIn('"data":', serialized) self.assertEqual(message["raw_fields"]["content"]["sha256"], hashlib.sha256("你好".encode()).hexdigest()) self.assertEqual(message["raw_fields"]["server_id"], "s1") def test_media_references_do_not_become_text(self): row = dict(self.row, content_type=3, content=b"*1*" + b"opaque-media-token-" * 12) message = self.normalize(row) self.assertEqual(message["content"], "[图片]") self.assertNotIn("opaque-media-token", json.dumps(message)) def test_transcribed_voice_wins_over_typed_placeholder(self): message = self.normalize(dict(self.row, content_type=4, content=b"voice-proto"), {(self.account, "s1"): "请问几点开始"}) self.assertEqual(message["content"], "请问几点开始") self.assertEqual(message["content_parse_status"], "decoded") def test_legacy_adapter_also_emits_plaintext_and_digest_only(self): adapter = WxworkArchiveAdapter(None, self.root, keys_map={}) message = adapter._normalized_message(self.account, self.row, {}, {}) self.assertEqual(message["content"], "1") self.assertEqual(message["raw_fields"]["content"]["encoding"], "local-body-digest") self.assertNotIn('"data":', json.dumps(message)) def database(self): path = self.root / "message.db" with closing(sqlite3.connect(path)) as db: db.execute("CREATE TABLE message_table(send_time REAL,conversation_id TEXT,sender_id TEXT,content BLOB,extra_content BLOB,content_type INTEGER,server_id TEXT,client_id TEXT)") for index in range(1, 6): db.execute("INSERT INTO message_table VALUES(?,?,?,?,?,?,?,?)", (1000 if index < 5 else 1001, "M:20002", "20002", self.row["content"], b"metadata-private", 2, "s" + str(index), "c" + str(index))) db.commit() return path def test_repair_is_bounded_resumable_and_never_rewinds_cloud_cursor(self): path = self.database() self.api.cursor = {"send_time": 1000, "rowid": 4} decrypted = [(str(path), "message.db", self.account)] first = self.importer.run(decrypted) self.assertEqual(first["content_repair_scanned"], 2) self.assertEqual(first["content_repair_pending_accounts"], 1) self.assertEqual(self.api.cursor, {"send_time": 1001.0, "rowid": 5}) second = self.importer.run(decrypted) self.assertEqual(second["content_repair_scanned"], 2) self.assertEqual(second["content_repair_pending_accounts"], 0) self.assertEqual(self.api.cursor, {"send_time": 1001.0, "rowid": 5}) self.assertEqual(self.importer.run(decrypted)["content_repaired"], 0) repairs = [payload for payload in self.api.payloads if "checkpoint" not in payload] self.assertEqual([item["messages"][0]["source_message_id"] for item in repairs], ["s1", "s2", "s3", "s4"]) self.assertTrue(all(item["messages"][0]["content"] == "1" for item in repairs)) self.assertTrue(all("_local_media_paths" not in item["messages"][0] for item in repairs)) def test_repair_failure_does_not_advance_its_local_checkpoint(self): path = self.database() self.api.cursor = {"send_time": 1000, "rowid": 4} original = self.api.import_messages def fail_repair(payload): if "checkpoint" not in payload: raise RuntimeError("synthetic network failure") return original(payload) self.api.import_messages = fail_repair with self.assertRaisesRegex(RuntimeError, "synthetic network"): self.importer.run([(str(path), "message.db", self.account)]) saved = json.loads(next((self.config.work_root / "content_repair").glob("*.json")).read_text()) self.assertEqual(saved["cursor"], {"send_time": 0, "rowid": 0}) self.assertEqual(self.api.cursor["rowid"], 5) self.api.import_messages = original result = self.importer.run([(str(path), "message.db", self.account)]) self.assertEqual(result["content_repaired"], 2) self.assertEqual(self.api.payloads[-2]["messages"][0]["source_message_id"], "s1") def test_repair_scope_separates_server_desktop_account_and_wecom_account(self): first, _ = self.importer._content_repair_state(self.account, {}) second, _ = self.importer._content_repair_state("other", {}) self.api.archive_scope = "different-desktop-account" third, _ = self.importer._content_repair_state(self.account, {}) self.api.base_url = "https://different.test" fourth, _ = self.importer._content_repair_state(self.account, {}) self.assertEqual(len({first, second, third, fourth}), 4) def test_repair_skips_clean_text_and_preserves_structured_file_cache_label(self): path = self.database() with closing(sqlite3.connect(path)) as db: db.execute("UPDATE message_table SET content=? WHERE rowid=1", (b"plain old text",)) db.commit() self.api.cursor = {"send_time": 1000, "rowid": 2} state_path, state = self.importer._content_repair_state(self.account, self.api.cursor) original = self.importer._normalize def with_file(account, row, users, conversations, voice): result = original(account, row, users, conversations, voice) if row["__archive_rowid"] == 2: result["attachment_metadata"] = [{"original_filename": "example.pdf", "status": "source_not_cached"}] return result with closing(sqlite3.connect(path)) as db, mock.patch.object(self.importer, "_normalize", side_effect=with_file): result = self.importer._repair_content(db, self.account, state_path, state, {}, {}, {}) self.assertEqual(result, {"scanned": 2, "reparsed": 0, "pending": False}) self.assertEqual(self.api.payloads, []) def test_repair_preserves_archived_voice_transcript_when_local_cache_is_missing(self): path = self.database() with closing(sqlite3.connect(path)) as local: local.execute("UPDATE message_table SET content_type=4 WHERE rowid=1") local.execute("UPDATE message_table SET content_type=9999 WHERE rowid=2") local.commit() db = admin_backend.Database(self.root / "voice-backend.db") db.initialize("TestPassword123!") store = ArchiveStore(db) store.initialize() voice = self.normalize(dict(self.row, content_type=4), {(self.account, "s1"): "原来的语音文字"}) voice.pop("_local_media_paths") store.import_messages({"source_account": {"external_account_id": self.account}, "messages": [voice]}, None, "test") self.api.import_messages = lambda payload: store.import_messages(payload, None, "test") state_path, state = self.importer._content_repair_state(self.account, {"send_time": 1000, "rowid": 2}) with closing(sqlite3.connect(path)) as local: result = self.importer._repair_content(local, self.account, state_path, state, {}, {}, {}) self.assertEqual(result["reparsed"], 0) with db.connect() as connection: self.assertEqual(connection.execute("SELECT content FROM archive_message").fetchone()[0], "原来的语音文字") self.assertEqual(connection.execute("SELECT COUNT(*) FROM archive_message_version").fetchone()[0], 1) # A real new transcription can still correct the historical body. state["complete"] = False state["cursor"] = {"send_time": 0.0, "rowid": 0} with closing(sqlite3.connect(path)) as local: result = self.importer._repair_content(local, self.account, state_path, state, {}, {}, {(self.account, "s1"): "已恢复的语音文字"}) self.assertEqual(result["reparsed"], 1) with db.connect() as connection: self.assertEqual(connection.execute("SELECT content FROM archive_message").fetchone()[0], "已恢复的语音文字") self.assertEqual(connection.execute("SELECT COUNT(*) FROM archive_message_version").fetchone()[0], 2) def test_existing_backend_upsert_repairs_one_message_without_checkpoint_change(self): db = admin_backend.Database(self.root / "backend.db") db.initialize("TestPassword123!") store = ArchiveStore(db) store.initialize() message = self.normalize() message.pop("_local_media_paths") message["content"] = "old-unparsed-body" payload = {"source_account": {"external_account_id": self.account}, "messages": [message], "checkpoint": {"send_time": 9999, "rowid": 55}} store.import_messages(payload, None, "test") repaired = dict(message, content="1") repair_payload = {"source_account": payload["source_account"], "messages": [repaired]} store.import_messages(repair_payload, None, "test") store.import_messages(repair_payload, None, "test") with db.connect() as connection: self.assertEqual(connection.execute("SELECT COUNT(*) FROM archive_message").fetchone()[0], 1) self.assertEqual(connection.execute("SELECT content FROM archive_message").fetchone()[0], "1") self.assertEqual(connection.execute("SELECT COUNT(*) FROM archive_message_version").fetchone()[0], 2) self.assertEqual(store.source_checkpoint(self.account, "message_table"), {"send_time": 9999, "rowid": 55}) if __name__ == "__main__": unittest.main()