Files
kefu/wechat_rpa/test_archive_api.py
T
2026-09-21 10:34:06 +08:00

1210 lines
49 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",
zyt_token_verifier=self.verify_zyt,
zyt_patient_searcher=self.search_zyt_patients,
)
)
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 verify_zyt(token: str) -> dict:
user_id = "202" if token == "zyt-user-2" else "101"
permissions = [] if token == "zyt-no-patient" else ["tcm.diagnosis/lists"]
return {
"user_id": user_id,
"sn": f"ZYT{user_id}",
"nickname": f"账号{user_id}",
"mobile": "13800138000",
"avatar": "",
"terminal": 7,
"status": "active",
"root": False,
"role_ids": [2],
"permissions": permissions,
"permissions_known": True,
}
@staticmethod
def search_zyt_patients(
token: str, keyword: str, page_no: int, page_size: int
) -> dict:
if not token.startswith("zyt-user-"):
raise RuntimeError("token invalid")
return {
"items": [
{
"patient_id": 9001,
"diagnosis_id": 9101,
"patient_name": f"{keyword}患者",
"phone_masked": "138****8000",
"gender": 1,
"age": 36,
"source_update_time": "2026-09-01 10:00:00",
}
],
"total": 1,
"page_no": page_no,
"page_size": page_size,
}
def desktop_login(self, token="zyt-user-1", device_id="desktop-device-01"):
response = self.client.post(
"/api/v2/desktop/auth/exchange",
headers={"Authorization": f"Bearer {token}"},
json={"device_id": device_id, "device_name": "test", "app_version": "1"},
)
self.assertEqual(response.status_code, 200, response.text)
return {
"headers": {"Authorization": f"Bearer {response.json()['access_token']}"},
"account": response.json()["account"],
}
@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_import_rejects_system_conversations_from_older_clients(self):
payload = self.sample_payload()
base = payload["messages"][0]
payload["messages"] = [
{
**base,
"source_message_id": f"system-{index}",
"server_id": f"system-server-{index}",
"conversation": {
"external_id": external_id,
"name": name,
},
}
for index, (external_id, name) in enumerate(
(
("O:5629501326797629", "服务 5629501326797629"),
("APPROVAL", "审批"),
("S:system-team", "企业微信团队"),
("S:external-user-01", "客户张三"),
),
start=1,
)
]
imported = self.client.post(
"/api/v2/archive/imports/messages",
headers=self.admin_headers,
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
self.assertEqual(imported.json()["received"], 4)
self.assertEqual(imported.json()["inserted"], 1)
self.assertEqual(imported.json()["skipped"], 3)
conversations = self.client.get(
"/api/v2/archive/conversations", headers=self.admin_headers
).json()["items"]
self.assertEqual(
[(row["external_id"], row["name"]) for row in conversations],
[("S:external-user-01", "客户张三")],
)
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_login_and_tenant_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)
session = self.desktop_login()
headers = session["headers"]
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"], [])
# 管理端必须明确选择这个桌面账号;默认租户里看不到它的数据。
legacy = local.get("/api/v2/archive/stats", headers=self.admin_headers).json()
scoped = local.get(
"/api/v2/archive/stats",
headers=self.admin_headers,
params={"account_id": session["account"]["id"]},
).json()
self.assertEqual(legacy["messages"], 0)
self.assertEqual(scoped["messages"], 1)
finally:
local.close()
def test_two_zyt_accounts_keep_identical_wecom_ids_separate(self):
first = self.desktop_login("zyt-user-1", "desktop-device-01")
imported_first = self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=first["headers"],
json=self.sample_payload(),
)
second = self.desktop_login("zyt-user-2", "desktop-device-02")
imported_second = self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=second["headers"],
json=self.sample_payload(),
)
self.assertEqual(imported_first.json()["inserted"], 1)
self.assertEqual(imported_second.json()["inserted"], 1)
for session in (first, second):
stats = self.client.get(
"/api/v2/archive/stats",
headers=self.admin_headers,
params={"account_id": session["account"]["id"]},
)
self.assertEqual(stats.status_code, 200, stats.text)
self.assertEqual(stats.json()["messages"], 1)
def test_all_accounts_aggregates_every_tenant(self):
"""账号下拉框选「全部」时,看到的是所有账号的数据,而不是某一个的。"""
first = self.desktop_login("zyt-user-1", "desktop-device-01")
self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=first["headers"],
json=self.sample_payload(),
)
second = self.desktop_login("zyt-user-2", "desktop-device-02")
self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=second["headers"],
json=self.sample_payload(),
)
stats = self.client.get(
"/api/v2/archive/stats",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(stats.status_code, 200, stats.text)
# 两个账号各一条,选「全部」就该是两条;只选一个账号时是一条
self.assertEqual(stats.json()["messages"], 2)
self.assertEqual(stats.json()["conversations"], 2)
conversations = self.client.get(
"/api/v2/archive/conversations",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(conversations.status_code, 200, conversations.text)
self.assertEqual(len(conversations.json()["items"]), 2)
people = self.client.get(
"/api/v2/archive/people",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(people.status_code, 200, people.text)
self.assertEqual(len(people.json()["items"]), 2)
def test_all_accounts_does_not_leak_past_an_admins_scope(self):
"""被限制到某几个账号的管理员,选「全部」只能看到这几个账号。"""
first = self.desktop_login("zyt-user-1", "desktop-device-01")
self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=first["headers"],
json=self.sample_payload(),
)
second = self.desktop_login("zyt-user-2", "desktop-device-02")
self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=second["headers"],
json=self.sample_payload(),
)
# 这个角色故意不给 desktop-account:all——有那个权限的人本来就该看到全部,
# 这条用例要验的是"被限制的人选全部时看到的是自己那几个账号"。
role = self.client.post(
"/api/v2/roles",
headers=self.admin_headers,
json={
"code": "scoped-admin",
"name": "受限管理员",
"permissions": ["im:read", "im:content:read", "stats:read"],
},
)
self.assertEqual(role.status_code, 200, role.text)
created = self.client.post(
"/api/v2/users",
headers=self.admin_headers,
json={
"username": "scoped",
"password": ADMIN_PASSWORD,
"role": "scoped-admin",
},
)
self.assertEqual(created.status_code, 200, created.text)
row = self.database.authenticate("scoped", ADMIN_PASSWORD)
self.database.change_password(
int(row["id"]), ADMIN_PASSWORD, ADMIN_PASSWORD + "x", "127.0.0.1"
)
# 把这个管理员限制到第一个账号
self.database.set_desktop_account_admins(
int(first["account"]["id"]), [int(row["id"])], 1, "127.0.0.1"
)
scoped_headers = self.login("scoped", ADMIN_PASSWORD + "x")
stats = self.client.get(
"/api/v2/archive/stats",
headers=scoped_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(stats.status_code, 200, stats.text)
self.assertEqual(stats.json()["messages"], 1, "越权看到了别的账号的消息")
denied = self.client.get(
"/api/v2/archive/stats",
headers=scoped_headers,
params={"account_id": second["account"]["id"]},
)
self.assertEqual(denied.status_code, 403)
def test_writes_still_require_one_account(self):
"""看数据可以跨账号,写不行——导出和绑定必须落到某一个账号名下。"""
session = self.desktop_login("zyt-user-1", "desktop-device-01")
self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=session["headers"],
json=self.sample_payload(),
)
export = self.client.post(
"/api/v2/archive/exports",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
json={"formats": ["sql"], "filters": {}},
)
self.assertEqual(export.status_code, 400, export.text)
self.assertIn("客户端账号", export.json()["detail"])
people = self.client.get(
"/api/v2/archive/people",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
).json()["items"]
bind = self.client.post(
f"/api/v2/archive/people/{people[0]['id']}/identities",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
json={
"identity_type": "wecom_userid",
"scope_id": "corp-001",
"external_id": "another-id",
},
)
self.assertEqual(bind.status_code, 400, bind.text)
self.assertIn("客户端账号", bind.json()["detail"])
def test_call_stats_can_cover_every_tenant(self):
first = self.desktop_login("zyt-user-1", "desktop-device-01")
second = self.desktop_login("zyt-user-2", "desktop-device-02")
for session in (first, second):
self.database.log_model_call({
"tenant_id": session["account"]["tenant_id"],
"purpose": "chat",
"chosen": "model-a",
"customer_text": "血糖13",
"reply_text": "空腹13偏高了",
})
everything = self.client.get(
"/api/v2/stats/model-calls",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(everything.status_code, 200, everything.text)
self.assertEqual(everything.json()["total"], 2)
single = self.client.get(
"/api/v2/stats/model-calls",
headers=self.admin_headers,
params={"account_id": first["account"]["id"]},
)
self.assertEqual(single.json()["total"], 1)
def test_tenant_scope_translates_the_three_cases(self):
"""范围翻译只有这一处,散到各个查询里迟早会漏掉一个。"""
from archive_store import ALL_TENANTS, tenant_scope
everything = tenant_scope(ALL_TENANTS)
self.assertEqual(everything.clause(), "1=1")
self.assertEqual(everything.params, [])
with self.assertRaises(ValueError):
everything.single
one = tenant_scope("zyt-abc")
self.assertEqual(one.clause("c"), "c.tenant_id=?")
self.assertEqual(one.params, ["zyt-abc"])
self.assertEqual(one.single, "zyt-abc")
several = tenant_scope(["zyt-a", "zyt-b", "zyt-a"])
self.assertEqual(several.clause("m"), "m.tenant_id IN (?,?)")
self.assertEqual(several.params, ["zyt-a", "zyt-b"])
with self.assertRaises(ValueError):
several.single
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.exe(155.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_patient_search_binding_and_unbind_stay_in_local_archive(self):
session = self.desktop_login()
account_id = session["account"]["id"]
imported = self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=session["headers"],
json=self.sample_payload(),
)
self.assertEqual(imported.status_code, 200, imported.text)
person = self.client.get(
"/api/v2/archive/people",
headers=self.admin_headers,
params={"account_id": account_id},
).json()["items"][0]
searched = self.client.get(
"/api/v2/archive/patients/search",
headers=self.admin_headers,
params={"account_id": account_id, "keyword": "张三"},
)
self.assertEqual(searched.status_code, 200, searched.text)
candidate = searched.json()["items"][0]
self.assertEqual(candidate["patient_id"], 9001)
self.assertEqual(candidate["phone_masked"], "138****8000")
bound = self.client.put(
f"/api/v2/archive/people/{person['id']}/patient-bindings",
headers=self.admin_headers,
params={"account_id": account_id},
json={
"patient_id": candidate["patient_id"],
"diagnosis_id": candidate["diagnosis_id"],
"relation_type": "self",
},
)
self.assertEqual(bound.status_code, 200, bound.text)
binding = bound.json()["binding"]
self.assertEqual(binding["patient_name"], "张三患者")
self.assertTrue(binding["is_primary"])
listed = self.client.get(
"/api/v2/archive/people",
headers=self.admin_headers,
params={"account_id": account_id},
).json()["items"][0]
self.assertEqual(listed["patient_binding_count"], 1)
self.assertEqual(listed["primary_patient_id"], 9001)
unbound = self.client.delete(
f"/api/v2/archive/people/{person['id']}/patient-bindings/{binding['id']}",
headers=self.admin_headers,
params={"account_id": account_id},
)
self.assertEqual(unbound.status_code, 200, unbound.text)
self.assertEqual(unbound.json()["person"]["patient_bindings"], [])
with self.database.connect() as db:
encrypted = db.execute(
"SELECT token_enc FROM archive_zyt_session"
).fetchone()[0]
actions = [
row[0]
for row in db.execute(
"SELECT action FROM archive_patient_binding_audit ORDER BY created_at"
).fetchall()
]
self.assertNotEqual(encrypted, "zyt-user-1")
self.assertEqual(actions, ["bind", "unbind"])
def test_desktop_can_bind_current_wecom_conversation_to_zyt_patient(self):
session = self.desktop_login()
payload = self.sample_payload()
imported = self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=session["headers"],
json=payload,
)
self.assertEqual(imported.status_code, 200, imported.text)
params = {
"external_account_id": payload["source_account"]["external_account_id"],
"conversation_external_id": payload["messages"][0]["conversation"][
"external_id"
],
}
context = self.client.get(
"/api/v2/archive/desktop/patient-context",
headers=session["headers"],
params=params,
)
self.assertEqual(context.status_code, 200, context.text)
self.assertTrue(context.json()["available"])
self.assertTrue(context.json()["permissions"]["can_bind"])
self.assertEqual(context.json()["person"]["display_name"], "张三")
searched = self.client.get(
"/api/v2/archive/desktop/patients/search",
headers=session["headers"],
params={"keyword": "张三"},
)
self.assertEqual(searched.status_code, 200, searched.text)
patient = searched.json()["items"][0]
bound = self.client.put(
"/api/v2/archive/desktop/patient-binding",
headers=session["headers"],
json={
**params,
"person_id": context.json()["person"]["id"],
"patient_id": patient["patient_id"],
"diagnosis_id": patient["diagnosis_id"],
"relation_type": "self",
},
)
self.assertEqual(bound.status_code, 200, bound.text)
binding = bound.json()["binding"]
self.assertEqual(binding["patient_name"], "张三患者")
self.assertEqual(
bound.json()["context"]["patient_bindings"][0]["patient_id"], 9001
)
unbound = self.client.delete(
"/api/v2/archive/desktop/patient-binding",
headers=session["headers"],
params={
**params,
"person_id": context.json()["person"]["id"],
"binding_id": binding["id"],
},
)
self.assertEqual(unbound.status_code, 200, unbound.text)
self.assertEqual(unbound.json()["context"]["patient_bindings"], [])
def test_desktop_patient_actions_follow_the_logged_in_zyt_permissions(self):
session = self.desktop_login(token="zyt-no-patient")
params = {
"external_account_id": "1688856770803435",
"conversation_external_id": "S:restricted-patient",
}
context = self.client.get(
"/api/v2/archive/desktop/patient-context",
headers=session["headers"],
params=params,
)
self.assertEqual(context.status_code, 200, context.text)
self.assertFalse(context.json()["available"])
self.assertFalse(context.json()["permissions"]["can_search"])
self.assertEqual(context.json()["people"], [])
searched = self.client.get(
"/api/v2/archive/desktop/patients/search",
headers=session["headers"],
params={"keyword": "张三"},
)
self.assertEqual(searched.status_code, 403, searched.text)
bound = self.client.put(
"/api/v2/archive/desktop/patient-binding",
headers=session["headers"],
json={
**params,
"patient_id": 9001,
"diagnosis_id": 9101,
"relation_type": "self",
},
)
self.assertEqual(bound.status_code, 403, bound.text)
self.assertIn("没有患者列表权限", bound.json()["detail"])
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_finished_exports_stay_readable_under_all_accounts(self):
"""导出要选具体账号,但导出**结果**是读,选「全部」时照样看得到、下得动。
这条是补的:`export_file` 当初漏了没跟着改租户范围,选「全部」点下载会
撞上一句"租户只能包含字母、数字、下划线和连字符"——那句话对着界面上的
下拉框根本看不懂。
"""
session = self.desktop_login("zyt-user-1", "desktop-device-01")
self.client.post(
"/api/v2/archive/desktop/imports/messages",
headers=session["headers"],
json=self.sample_payload(),
)
created = self.client.post(
"/api/v2/archive/exports",
headers=self.admin_headers,
params={"account_id": session["account"]["id"]},
json={"formats": ["sql"], "filters": {}},
)
self.assertEqual(created.status_code, 200, created.text)
job_id = created.json()["job"]["id"]
listed = self.client.get(
"/api/v2/archive/exports",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(listed.status_code, 200, listed.text)
self.assertIn(job_id, [item["id"] for item in listed.json()["jobs"]])
job = self.client.get(
f"/api/v2/archive/exports/{job_id}",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(job.status_code, 200, job.text)
sql_file = next(
item for item in job.json()["job"]["files"] if item["format"] == "sql"
)
downloaded = self.client.get(
f"/api/v2/archive/export-files/{sql_file['id']}/download",
headers=self.admin_headers,
params={"account_id": admin_api.ALL_ACCOUNTS},
)
self.assertEqual(downloaded.status_code, 200, downloaded.text)
self.assertIn(b"archive_messages", 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()