210 lines
12 KiB
Python
210 lines
12 KiB
Python
# -*- 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()
|