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

675 lines
28 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""聊天归档新模块的接口、去重、导出和密钥保护回归测试。"""
import sqlite3
import tempfile
from pathlib import Path
from unittest import TestCase, main, mock
from fastapi.testclient import TestClient
import admin_api
import admin_backend as backend
from archive_store import MULTIPART_THRESHOLD_BYTES
from archive_wxwork_adapter import WxworkArchiveAdapter
ADMIN_PASSWORD = "Archive@2026Admin"
VIEWER_PASSWORD = "Archive@2026Viewer"
class ArchiveApiTest(TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.database = backend.Database(self.root / "archive.db")
self.database.initialize("Admin@123456")
admin = self.database.authenticate("admin", "Admin@123456")
self.database.change_password(
int(admin["id"]), "Admin@123456", ADMIN_PASSWORD, "127.0.0.1"
)
self.client = TestClient(admin_api.create_app(self.root / "archive.db"))
self.admin_headers = self.login("admin", ADMIN_PASSWORD)
def tearDown(self):
self.client.close()
self.temp.cleanup()
def login(self, username: str, password: str) -> dict[str, str]:
response = self.client.post(
"/api/v2/auth/login",
json={"username": username, "password": password},
)
self.assertEqual(response.status_code, 200, response.text)
return {"Authorization": f"Bearer {response.json()['access_token']}"}
@staticmethod
def sample_payload() -> dict:
return {
"source_account": {
"external_account_id": "wxwork-account-01",
"display_name": "客服一号",
"corp_scope_id": "corp-001",
},
"source_table": "message_table",
"checkpoint": {"rowid": 1001},
"messages": [
{
"source_message_id": "source-message-001",
"server_id": "server-message-001",
"conversation": {
"external_id": "S:external-user-01",
"name": "客户张三",
},
"sender": {
"external_id": "zhangsan",
"display_name": "张三",
"identity_type": "wecom_userid",
"scope_id": "corp-001",
"verified": True,
},
"sent_at": 1_724_472_000_000,
"message_type": "text",
"direction": "inbound",
"content": "你好,这是一条归档消息",
}
],
}
def test_import_is_idempotent_and_query_is_structured(self):
first = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=self.sample_payload(),
)
self.assertEqual(first.status_code, 200, first.text)
self.assertEqual(first.json()["inserted"], 1)
second = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=self.sample_payload(),
)
self.assertEqual(second.status_code, 200, second.text)
self.assertEqual(second.json()["inserted"], 0)
self.assertEqual(second.json()["duplicates"], 1)
stats = self.client.get(
"/api/v2/archive/stats", headers=self.admin_headers
).json()
self.assertEqual(stats["messages"], 1)
self.assertEqual(stats["people"], 1)
self.assertEqual(stats["conversations"], 1)
conversations = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()
self.assertEqual(conversations["items"][0]["message_count"], 1)
conversation_id = conversations["items"][0]["id"]
messages = self.client.get(
f"/api/v2/archive/conversations/{conversation_id}/messages",
headers=self.admin_headers,
).json()
self.assertEqual(messages["items"][0]["sender_name"], "张三")
with self.database.connect() as db:
member_count = db.execute(
"SELECT COUNT(*) FROM archive_conversation_member"
).fetchone()[0]
self.assertEqual(member_count, 1)
def test_messages_are_returned_in_chronological_order(self):
payload = self.sample_payload()
newer = payload["messages"][0]
newer["content"] = "较新的消息"
older = {
**newer,
"source_message_id": "source-message-older",
"server_id": "server-message-older",
"sent_at": newer["sent_at"] - 60_000,
"content": "较早的消息",
}
payload["messages"] = [newer, older]
imported = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
conversation_id = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()["items"][0]["id"]
items = self.client.get(
f"/api/v2/archive/conversations/{conversation_id}/messages",
headers=self.admin_headers,
).json()["items"]
self.assertEqual(
[item["content"] for item in items], ["较早的消息", "较新的消息"]
)
def test_metadata_refresh_replaces_ids_with_nicknames_without_new_messages(self):
payload = self.sample_payload()
payload["messages"][0]["conversation"]["name"] = "S:external-user-01"
payload["messages"][0]["sender"]["display_name"] = "zhangsan"
imported = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
synced = self.client.post(
"/api/v2/archive/imports/metadata",
headers=self.admin_headers,
json={
"source_account": payload["source_account"],
"people": [
{
"external_id": "zhangsan",
"display_name": "张三昵称",
"identity_type": "wecom_userid",
"scope_id": "corp-001",
}
],
"conversations": [
{
"external_id": "S:external-user-01",
"name": "张三昵称",
}
],
},
)
self.assertEqual(synced.status_code, 200, synced.text)
self.assertEqual(synced.json()["conversations_updated"], 1)
conversations = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()
self.assertEqual(conversations["items"][0]["name"], "张三昵称")
conversation_id = conversations["items"][0]["id"]
messages = self.client.get(
f"/api/v2/archive/conversations/{conversation_id}/messages",
headers=self.admin_headers,
).json()
self.assertEqual(messages["items"][0]["sender_name"], "张三昵称")
def test_short_nested_protobuf_text_is_normalized_but_raw_event_is_preserved(self):
payload = self.sample_payload()
payload["messages"][0]["content"] = "0a07080012030a0131"
imported = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
conversation = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()["items"][0]
self.assertEqual(conversation["last_content"], "1")
message = self.client.get(
f"/api/v2/archive/conversations/{conversation['id']}/messages",
headers=self.admin_headers,
).json()["items"][0]
self.assertEqual(message["content"], "1")
with self.database.connect() as db:
raw_payload = db.execute(
"SELECT payload_json FROM archive_raw_event LIMIT 1"
).fetchone()[0]
stored = db.execute("SELECT content FROM archive_message").fetchone()[0]
self.assertIn("0a07080012030a0131", raw_payload)
self.assertEqual(stored, "1")
def test_message_returns_attachment_details_and_hides_binary_payload(self):
with self.database.connect() as db:
db.execute(
"""INSERT INTO archive_media_object
(id,tenant_id,bucket,region,object_key,sha256,size_bytes,mime_type,
original_filename,media_type,status,created_at,verified_at)
VALUES ('media-image','default','bucket-1','ap-guangzhou','image/1',
?,321,'image/png','截图.png','image','ready',?,?)""",
("c" * 64, "2026-08-24T00:00:00+00:00", "2026-08-24T00:00:00+00:00"),
)
db.commit()
payload = self.sample_payload()
payload["messages"][0].update(
{
"message_type": "截图",
"content": "30" * 100,
"media_ids": ["media-image"],
}
)
imported = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
conversations = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()
conversation_id = conversations["items"][0]["id"]
self.assertEqual(conversations["items"][0]["last_content"], "[截图]")
message = self.client.get(
f"/api/v2/archive/conversations/{conversation_id}/messages",
headers=self.admin_headers,
).json()["items"][0]
self.assertEqual(message["content"], "[截图]")
self.assertEqual(message["attachment_count"], 1)
self.assertEqual(message["attachments"][0]["original_filename"], "截图.png")
self.assertEqual(message["attachments"][0]["media_type"], "image")
cos = mock.Mock()
cos.get_presigned_download_url.return_value = "https://cos.example/read-image"
store = self.client.app.state.archive_store
with mock.patch.object(store, "_cos_client", return_value=({}, cos)):
access = self.client.post(
"/api/v2/archive/media/access-urls",
headers=self.admin_headers,
json={"media_ids": ["media-image"], "expires": 600},
)
self.assertEqual(access.status_code, 200, access.text)
self.assertEqual(access.json()["items"][0]["url"], "https://cos.example/read-image")
def test_desktop_auto_backup_endpoint_is_loopback_and_key_protected(self):
local = TestClient(
self.client.app, client=("127.0.0.1", 50123)
)
try:
denied = local.post(
"/api/v2/archive/desktop/imports/messages",
headers={"X-Desktop-Sync-Key": "wrong"},
json=self.sample_payload(),
)
self.assertEqual(denied.status_code, 401, denied.text)
headers = {"X-Desktop-Sync-Key": backend.DESKTOP_SYNC_KEY}
imported = local.post(
"/api/v2/archive/desktop/imports/messages",
headers=headers,
json=self.sample_payload(),
)
self.assertEqual(imported.status_code, 200, imported.text)
self.assertEqual(imported.json()["inserted"], 1)
checkpoint = local.get(
"/api/v2/archive/desktop/checkpoint",
headers=headers,
params={"external_account_id": "wxwork-account-01"},
)
self.assertEqual(checkpoint.status_code, 200, checkpoint.text)
self.assertEqual(checkpoint.json()["checkpoint"], {"rowid": 1001})
advanced = local.post(
"/api/v2/archive/desktop/checkpoint",
headers=headers,
json={
"source_account": self.sample_payload()["source_account"],
"source_table": "message_table",
"checkpoint": {"send_time": 2000, "rowid": 1002},
},
)
self.assertEqual(advanced.status_code, 200, advanced.text)
self.assertEqual(
advanced.json()["checkpoint"], {"send_time": 2000, "rowid": 1002}
)
pending = local.get(
"/api/v2/archive/desktop/pending-attachments",
headers=headers,
params={"external_account_id": "wxwork-account-01"},
)
self.assertEqual(pending.status_code, 200, pending.text)
self.assertEqual(pending.json()["source_message_ids"], [])
finally:
local.close()
def test_uncached_file_is_queryable_and_queued_for_desktop_retry(self):
payload = self.sample_payload()
message = payload["messages"][0]
message["message_type"] = "文件"
message["content"] = "[文件] installer.exe155.32 MB,源文件未缓存)"
message["attachment_metadata"] = [
{
"original_filename": "installer.exe",
"size_bytes": 162_862_195,
"checksum": "9D4295ED638B290925A5A960D644934E",
"media_type": "file",
"status": "source_not_cached",
"source_reference_sha256": "a" * 64,
}
]
imported = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
conversations = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()
conversation_id = conversations["items"][0]["id"]
item = self.client.get(
f"/api/v2/archive/conversations/{conversation_id}/messages",
headers=self.admin_headers,
).json()["items"][0]
self.assertEqual(item["message_type"], "文件")
self.assertEqual(item["attachment_count"], 1)
self.assertEqual(item["attachments"][0]["status"], "source_not_cached")
self.assertEqual(item["attachments"][0]["original_filename"], "installer.exe")
self.assertEqual(
self.client.app.state.archive_store.pending_attachment_source_ids(
"wxwork-account-01"
),
["source-message-001"],
)
def test_changed_source_message_keeps_raw_and_normalized_versions(self):
payload = self.sample_payload()
self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
payload["messages"][0]["content"] = "这条消息已编辑"
changed = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(changed.status_code, 200, changed.text)
self.assertEqual(changed.json()["duplicates"], 1)
with self.database.connect() as db:
self.assertEqual(
db.execute("SELECT COUNT(*) FROM archive_raw_event").fetchone()[0], 2
)
self.assertEqual(
db.execute("SELECT COUNT(*) FROM archive_message_version").fetchone()[0],
2,
)
content = db.execute("SELECT content FROM archive_message").fetchone()[0]
self.assertEqual(content, "这条消息已编辑")
def test_wecom_identity_can_be_bound_to_a_person(self):
self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=self.sample_payload(),
)
person = self.client.get(
"/api/v2/archive/people", headers=self.admin_headers
).json()["items"][0]
response = self.client.post(
f"/api/v2/archive/people/{person['id']}/identities",
headers=self.admin_headers,
json={
"identity_type": "wecom_open_userid",
"scope_id": "corp-001",
"external_id": "open-user-001",
"verified": True,
},
)
self.assertEqual(response.status_code, 200, response.text)
identities = response.json()["person"]["identities"]
self.assertEqual(len(identities), 2)
self.assertIn("wecom_open_userid", {row["identity_type"] for row in identities})
def test_export_builds_sql_and_spreadsheet_files(self):
self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=self.sample_payload(),
)
response = self.client.post(
"/api/v2/archive/exports",
headers=self.admin_headers,
json={"formats": ["sql", "xlsx"], "filters": {}},
)
self.assertEqual(response.status_code, 200, response.text)
job_id = response.json()["job"]["id"]
job = self.client.get(
f"/api/v2/archive/exports/{job_id}", headers=self.admin_headers
).json()["job"]
self.assertEqual(job["status"], "completed", job.get("error_message"))
self.assertEqual(job["total_rows"], 1)
names = {item["file_name"] for item in job["files"]}
self.assertIn("archive.sql", names)
self.assertIn("archive.xlsx", names)
self.assertIn("archive_bundle.zip", names)
sql_file = next(item for item in job["files"] if item["format"] == "sql")
downloaded = self.client.get(
f"/api/v2/archive/export-files/{sql_file['id']}/download",
headers=self.admin_headers,
)
self.assertEqual(downloaded.status_code, 200, downloaded.text)
self.assertIn(b"archive_messages", downloaded.content)
self.assertIn(b"archive_people", downloaded.content)
self.assertIn(b"archive_pending_attachments", downloaded.content)
def test_cos_secrets_are_encrypted_and_never_echoed(self):
secret_id = "test-secret-id-that-must-not-echo"
secret_key = "test-secret-key-that-must-not-echo"
response = self.client.put(
"/api/v2/archive/storage",
headers=self.admin_headers,
json={
"bucket": "archive-1234567890",
"region": "ap-guangzhou",
"secret_id": secret_id,
"secret_key": secret_key,
"enabled": False,
},
)
self.assertEqual(response.status_code, 200, response.text)
body = response.text
self.assertNotIn(secret_id, body)
self.assertNotIn(secret_key, body)
self.assertTrue(response.json()["storage"]["secret_id_present"])
with self.database.connect() as db:
row = db.execute(
"SELECT secret_id_enc,secret_key_enc FROM archive_storage_config WHERE id=1"
).fetchone()
self.assertNotEqual(row["secret_id_enc"], secret_id)
self.assertNotEqual(row["secret_key_enc"], secret_key)
def test_media_is_direct_uploaded_and_verified_in_cos(self):
digest = "a" * 64
cos = mock.Mock()
cos.get_presigned_url.return_value = "https://cos.example/upload-signature"
cos.head_object.return_value = {
"Content-Length": "1234",
"x-cos-meta-sha256": digest,
"x-cos-hash-crc64ecma": "99887766",
"ETag": '"etag-value"',
"x-cos-version-id": "version-1",
}
config = {
"bucket": "archive-1234567890",
"region": "ap-guangzhou",
"media_prefix": "archive/media",
"encryption_mode": "AES256",
}
store = self.client.app.state.archive_store
with mock.patch.object(store, "_cos_client", return_value=(config, cos)):
prepared = self.client.post(
"/api/v2/archive/media/prepare",
headers=self.admin_headers,
json={
"sha256": digest,
"size_bytes": 1234,
"mime_type": "image/png",
"original_filename": "image.png",
},
)
self.assertEqual(prepared.status_code, 200, prepared.text)
body = prepared.json()
self.assertEqual(body["upload_url"], "https://cos.example/upload-signature")
self.assertEqual(body["required_headers"]["x-cos-meta-sha256"], digest)
self.assertIn(digest, body["media"]["object_key"])
completed = self.client.post(
f"/api/v2/archive/media/{body['media']['id']}/complete",
headers=self.admin_headers,
)
self.assertEqual(completed.status_code, 200, completed.text)
media = completed.json()["media"]
self.assertEqual(media["status"], "ready")
self.assertEqual(media["version_id"], "version-1")
self.assertEqual(media["crc64"], "99887766")
def test_large_video_uses_resumable_multipart_cos_upload(self):
digest = "b" * 64
size = MULTIPART_THRESHOLD_BYTES + 5 * 1024 * 1024
cos = mock.Mock()
cos.create_multipart_upload.return_value = {"UploadId": "upload-large-1"}
cos.list_parts.return_value = {"Part": [], "IsTruncated": "false"}
cos.get_presigned_url.side_effect = lambda **values: (
f"https://cos.example/part-{values['Params']['partNumber']}"
)
cos.complete_multipart_upload.return_value = {"ETag": '"multipart-etag"'}
head_without_metadata = {
"Content-Length": str(size),
"x-cos-hash-crc64ecma": "112233",
"ETag": '"multipart-etag"',
}
head_with_metadata = dict(
head_without_metadata, **{"x-cos-meta-sha256": digest}
)
cos.head_object.side_effect = [head_without_metadata, head_with_metadata]
cos.copy_object.return_value = {"ETag": '"multipart-etag"'}
config = {
"bucket": "archive-1234567890",
"region": "ap-guangzhou",
"media_prefix": "archive/media",
"encryption_mode": "AES256",
}
store = self.client.app.state.archive_store
with mock.patch.object(store, "_cos_client", return_value=(config, cos)):
prepared = self.client.post(
"/api/v2/archive/media/prepare",
headers=self.admin_headers,
json={
"sha256": digest,
"size_bytes": size,
"mime_type": "video/mp4",
"original_filename": "large-video.mp4",
},
)
self.assertEqual(prepared.status_code, 200, prepared.text)
body = prepared.json()
self.assertEqual(body["upload_mode"], "multipart")
self.assertGreater(len(body["multipart"]["parts"]), 1)
self.assertEqual(body["multipart"]["upload_id"], "upload-large-1")
parts = [
{"part_number": item["part_number"], "etag": f"etag-{item['part_number']}"}
for item in body["multipart"]["parts"]
]
completed = self.client.post(
f"/api/v2/archive/media/{body['media']['id']}/multipart-complete",
headers=self.admin_headers,
json={"upload_id": "upload-large-1", "parts": parts},
)
self.assertEqual(completed.status_code, 200, completed.text)
self.assertEqual(completed.json()["media"]["status"], "ready")
call = cos.complete_multipart_upload.call_args.kwargs
self.assertEqual(call["UploadId"], "upload-large-1")
self.assertEqual(len(call["MultipartUpload"]["Part"]), len(parts))
self.assertEqual(
cos.copy_object.call_args.kwargs["Metadata"],
{"x-cos-meta-sha256": digest},
)
with self.database.connect() as db:
remaining = db.execute(
"SELECT COUNT(*) FROM archive_media_upload_session"
).fetchone()[0]
self.assertEqual(remaining, 0)
def test_existing_wxwork_database_adapter_is_incremental(self):
account = "10001"
data_dir = self.root / "wxwork" / account / "Data"
data_dir.mkdir(parents=True)
with sqlite3.connect(data_dir / "message.db") as db:
db.execute(
"""CREATE TABLE message_table(
send_time REAL,conversation_id TEXT,sender_id TEXT,content BLOB,
content_type INTEGER,server_id TEXT,client_id TEXT)"""
)
db.executemany(
"INSERT INTO message_table VALUES (?,?,?,?,?,?,?)",
[
(1000, "R:group-1", "20002", "群聊消息", 2, "server-1", "client-1"),
(1000, "S:10001_20002", account, "自己发送", 2, "server-2", "client-2"),
],
)
db.close()
with sqlite3.connect(data_dir / "user.db") as db:
db.execute(
"CREATE TABLE user_table(id TEXT,name TEXT,real_name TEXT,account TEXT)"
)
db.executemany(
"INSERT INTO user_table VALUES (?,?,?,?)",
[(account, "客服账号", "", ""), ("20002", "张三", "", "")],
)
db.close()
with sqlite3.connect(data_dir / "session.db") as db:
db.execute(
"""CREATE TABLE conversation_table(
id TEXT,name TEXT,roomname_remark TEXT,session_id TEXT)"""
)
db.executemany(
"INSERT INTO conversation_table VALUES (?,?,?,?)",
[
("R:group-1", "项目群", "", ""),
("S:10001_20002", "张三", "", ""),
],
)
db.close()
adapter = WxworkArchiveAdapter(
self.client.app.state.archive_store,
self.root / "wxwork",
keys_map={},
cache_dir=self.root / "decrypt-cache",
)
admin_id = int(self.database.authenticate("admin", ADMIN_PASSWORD)["id"])
first = adapter.import_all(admin_id, batch_size=1)
second = adapter.import_all(admin_id, batch_size=1)
self.assertEqual(first["inserted"], 2)
self.assertEqual(first["batches"], 2)
self.assertEqual(second["inserted"], 0)
with self.database.connect() as db:
conversations = {
row[0] for row in db.execute("SELECT external_id FROM archive_conversation")
}
directions = {
row[0] for row in db.execute("SELECT direction FROM archive_message")
}
self.assertIn("R:group-1", conversations)
self.assertEqual(directions, {"inbound", "outbound"})
def test_viewer_has_no_archive_access_until_explicitly_granted(self):
admin_id = int(self.database.authenticate("admin", ADMIN_PASSWORD)["id"])
self.database.create_user(
"archive-viewer", VIEWER_PASSWORD, "viewer", admin_id, "127.0.0.1"
)
viewer = self.database.authenticate("archive-viewer", VIEWER_PASSWORD)
self.database.change_password(
int(viewer["id"]),
VIEWER_PASSWORD,
VIEWER_PASSWORD + "x",
"127.0.0.1",
)
headers = self.login("archive-viewer", VIEWER_PASSWORD + "x")
overview = self.client.get("/api/v2/archive/stats", headers=headers)
self.assertEqual(overview.status_code, 403)
self.assertIn("im:read", overview.json()["detail"])
denied = self.client.get(
"/api/v2/archive/conversations", headers=headers
)
self.assertEqual(denied.status_code, 403)
self.assertIn("im:content:read", denied.json()["detail"])
if __name__ == "__main__":
main()