test_business_assistant.py 22 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616
  1. from __future__ import annotations
  2. import json
  3. from datetime import UTC, datetime, timedelta
  4. from types import SimpleNamespace
  5. import pytest
  6. from aiohttp import CookieJar
  7. from aiohttp.test_utils import TestClient, TestServer
  8. def connection_payload(connection_id: str = "conn-a", *, can_reply: bool = True) -> dict:
  9. return {
  10. "id": connection_id,
  11. "user": {"id": 900, "username": "owner", "first_name": "店主"},
  12. "user_chat_id": 900,
  13. "date": int(datetime.now(UTC).timestamp()),
  14. "rights": {"can_reply": can_reply, "can_read_messages": True},
  15. "is_enabled": True,
  16. }
  17. def customer_message(
  18. *,
  19. message_id: int,
  20. text: str | None = "营业时间是什么?",
  21. chat_id: int = 501,
  22. sender_id: int = 501,
  23. sent_at: datetime | None = None,
  24. **values,
  25. ) -> dict:
  26. message = {
  27. "business_connection_id": "conn-a",
  28. "message_id": message_id,
  29. "date": int((sent_at or datetime.now(UTC)).timestamp()),
  30. "chat": {"id": chat_id, "type": "private"},
  31. "from": {"id": sender_id, "username": f"user{sender_id}", "first_name": "客户"},
  32. **values,
  33. }
  34. if text is not None:
  35. message["text"] = text
  36. return message
  37. class FakeProvider:
  38. configured = True
  39. def __init__(self, *, error: Exception | None = None) -> None:
  40. self.error = error
  41. self.calls: list[dict] = []
  42. async def decide(self, **values):
  43. self.calls.append(values)
  44. if self.error:
  45. raise self.error
  46. entry_id = str(values["knowledge"][0]["entry_id"])
  47. return {
  48. "action": "answer",
  49. "reply": "每天 09:00 至 18:00 营业。",
  50. "handoff_reason": "",
  51. "matched_entry_ids": [entry_id],
  52. "summary": "客户询问营业时间。",
  53. }
  54. async def test_connection(self):
  55. return {"ok": True, "model": "test-model", "response_id": "response-1"}
  56. class FakeBusinessApi:
  57. def __init__(self) -> None:
  58. self.sent: list[dict] = []
  59. self.read: list[tuple[str, int, int]] = []
  60. async def send_message(
  61. self,
  62. chat_id: int,
  63. text: str,
  64. *,
  65. business_connection_id: str = "",
  66. reply_markup: dict | None = None,
  67. ) -> dict:
  68. self.sent.append(
  69. {
  70. "chat_id": chat_id,
  71. "text": text,
  72. "business_connection_id": business_connection_id,
  73. "reply_markup": reply_markup,
  74. }
  75. )
  76. return {
  77. "message_id": 1000 + len(self.sent),
  78. "sender_business_bot": {"id": 999},
  79. }
  80. async def read_business_message(
  81. self, connection_id: str, chat_id: int, message_id: int
  82. ) -> None:
  83. self.read.append((connection_id, chat_id, message_id))
  84. async def prepare_runtime(app_modules, *, provider: FakeProvider | None = None):
  85. dbassistant = app_modules.load("wbb.utils.dbassistant")
  86. service = app_modules.load("wbb.services.business_assistant")
  87. await dbassistant.upsert_business_connection(connection_payload())
  88. await dbassistant.update_account_settings("conn-a", {"assistant_enabled": True})
  89. knowledge = await dbassistant.create_knowledge_entry(
  90. "conn-a",
  91. {
  92. "question": "营业时间是什么?",
  93. "aliases": ["几点营业"],
  94. "keywords": ["营业时间", "营业"],
  95. "answer": "每天 09:00 至 18:00 营业。",
  96. "priority": 10,
  97. },
  98. )
  99. selected_provider = provider or FakeProvider()
  100. runtime = service.BusinessAssistantRuntime(
  101. token="123456:" + "A" * 30,
  102. session=SimpleNamespace(),
  103. provider=selected_provider,
  104. )
  105. runtime.api = FakeBusinessApi()
  106. return dbassistant, service, runtime, knowledge
  107. async def test_connection_settings_knowledge_quota_and_bot_scope(app_modules):
  108. dbassistant = app_modules.load("wbb.utils.dbassistant")
  109. await dbassistant.upsert_business_connection(connection_payload("conn-a"))
  110. await dbassistant.upsert_business_connection(connection_payload("conn-b"))
  111. settings = await dbassistant.update_account_settings(
  112. "conn-a",
  113. {
  114. "assistant_enabled": True,
  115. "account_daily_limit": 2,
  116. "customer_daily_limit": 1,
  117. "timezone": "Asia/Shanghai",
  118. },
  119. )
  120. first = await dbassistant.create_knowledge_entry(
  121. "conn-a",
  122. {
  123. "question": "营业时间",
  124. "aliases": ["几点开门"],
  125. "keywords": ["开门"],
  126. "answer": "九点开门。",
  127. "priority": 20,
  128. },
  129. )
  130. await dbassistant.create_knowledge_entry(
  131. "conn-b",
  132. {"question": "营业时间", "keywords": ["开门"], "answer": "十点开门。"},
  133. )
  134. await app_modules.wbb.db.business_assistant_knowledge.insert_one(
  135. {
  136. "bot_id": "foreign-bot",
  137. "entry_id": "foreign-entry",
  138. "connection_id": "conn-a",
  139. "question": "营业时间",
  140. "aliases": [],
  141. "keywords": ["开门"],
  142. "answer": "不应跨 Bot 返回。",
  143. "priority": 1000,
  144. "enabled": True,
  145. "updated_at": datetime.now(UTC),
  146. }
  147. )
  148. matched = await dbassistant.match_knowledge("conn-a", "你们几点开门?")
  149. assert [item["entry_id"] for item in matched] == [first["entry_id"]]
  150. items, total = await dbassistant.list_knowledge_entries("conn-a")
  151. assert total == 1
  152. assert items[0]["answer"] == "九点开门。"
  153. await app_modules.wbb.db.business_assistant_connections.insert_one(
  154. {
  155. "bot_id": "foreign-bot",
  156. "connection_id": "foreign-connection",
  157. "updated_at": datetime.now(UTC),
  158. }
  159. )
  160. connections, total = await dbassistant.list_business_connections(page_size=100)
  161. assert total == 2
  162. assert {item["connection_id"] for item in connections} == {"conn-a", "conn-b"}
  163. assert await dbassistant.reserve_ai_usage("conn-a", 101, settings) == (True, "")
  164. assert await dbassistant.reserve_ai_usage("conn-a", 101, settings) == (
  165. False,
  166. "customer_daily_limit",
  167. )
  168. assert await dbassistant.reserve_ai_usage("conn-a", 102, settings) == (True, "")
  169. assert await dbassistant.reserve_ai_usage("conn-a", 103, settings) == (
  170. False,
  171. "account_daily_limit",
  172. )
  173. usage = await dbassistant.usage_metrics(connection_id="conn-a")
  174. assert usage["ai_calls"] == 2
  175. assert usage["customers"] == 2
  176. with pytest.raises(dbassistant.AssistantDataError) as error:
  177. dbassistant.normalize_account_settings({"timezone": "invalid/timezone"})
  178. assert error.value.code == "invalid_setting"
  179. def test_ai_contract_and_reply_window_validation(app_modules):
  180. service = app_modules.load("wbb.services.business_assistant")
  181. parsed = service.parse_ai_decision(
  182. json.dumps(
  183. {
  184. "action": "answer",
  185. "reply": "标准答案",
  186. "matched_entry_ids": ["entry-1", "foreign"],
  187. "summary": "摘要",
  188. }
  189. ),
  190. allowed_entry_ids={"entry-1"},
  191. )
  192. assert parsed["matched_entry_ids"] == ["entry-1"]
  193. with pytest.raises(service.AssistantProviderError):
  194. service.parse_ai_decision(
  195. '{"action":"answer","reply":"没有引用"}', allowed_entry_ids={"entry-1"}
  196. )
  197. now = datetime.now(UTC)
  198. assert service.business_reply_window_open(
  199. {"date": int((now - timedelta(hours=23)).timestamp())}, now=now
  200. )
  201. assert not service.business_reply_window_open(
  202. {"date": int((now - timedelta(hours=25)).timestamp())}, now=now
  203. )
  204. assert not service.business_reply_window_open({}, now=now)
  205. async def test_business_updates_are_idempotent_and_manual_reply_pauses(app_modules):
  206. dbassistant, _service, runtime, knowledge = await prepare_runtime(app_modules)
  207. update = {
  208. "update_id": 11,
  209. "business_message": customer_message(message_id=1),
  210. }
  211. await runtime.process_update(update)
  212. await runtime.process_update(update)
  213. assert len(runtime.api.sent) == 1
  214. assert runtime.api.sent[0]["business_connection_id"] == "conn-a"
  215. assert runtime.api.read == [("conn-a", 501, 1)]
  216. assert len(runtime.provider.calls) == 1
  217. assert runtime.provider.calls[0]["knowledge"][0]["entry_id"] == knowledge["entry_id"]
  218. conversation = await dbassistant.get_conversation_by_chat("conn-a", 501)
  219. assert conversation["summary"] == "客户询问营业时间。"
  220. messages = await dbassistant.recent_conversation_messages(
  221. conversation["conversation_id"], limit=10
  222. )
  223. assert [item["direction"] for item in messages] == ["incoming", "assistant"]
  224. assert messages[0]["expires_at"] - messages[0]["created_at"] == timedelta(days=30)
  225. message_indexes = await app_modules.wbb.db.business_assistant_messages.index_information()
  226. assert message_indexes["expires_at_1"]["expireAfterSeconds"] == 0
  227. await runtime.process_update(
  228. {
  229. "update_id": 12,
  230. "business_message": customer_message(
  231. message_id=2, text="账号本人回复", sender_id=900
  232. ),
  233. }
  234. )
  235. paused = await dbassistant.get_conversation(conversation["conversation_id"])
  236. assert paused["status"] == "human_paused"
  237. assert paused["handoff_reason"] == "human_reply_detected"
  238. assert paused["customer"]["id"] == 501
  239. await runtime.process_update(
  240. {
  241. "update_id": 13,
  242. "business_message": customer_message(
  243. message_id=3,
  244. sender_business_bot={"id": 999},
  245. ),
  246. }
  247. )
  248. await runtime.process_update(
  249. {
  250. "update_id": 14,
  251. "business_message": customer_message(
  252. message_id=4,
  253. is_from_offline=True,
  254. ),
  255. }
  256. )
  257. assert len(runtime.api.sent) == 1
  258. async def test_handoff_sensitive_non_text_provider_error_and_expired_window(app_modules):
  259. dbassistant, service, runtime, _knowledge = await prepare_runtime(app_modules)
  260. await dbassistant.update_account_settings(
  261. "conn-a",
  262. {
  263. "notification_destination": "both",
  264. "ops_group_id": -100123,
  265. "account_daily_limit": 1,
  266. "customer_daily_limit": 1,
  267. },
  268. )
  269. await runtime.process_update(
  270. {
  271. "update_id": 21,
  272. "business_message": customer_message(
  273. message_id=1, text="我要退款", chat_id=601, sender_id=601
  274. ),
  275. }
  276. )
  277. sensitive = await dbassistant.get_conversation_by_chat("conn-a", 601)
  278. assert sensitive["status"] == "handoff"
  279. assert sensitive["handoff_reason"] == "sensitive_request"
  280. assert runtime.api.sent[-1]["business_connection_id"] == ""
  281. assert "message_id=1" in runtime.api.sent[-1]["text"]
  282. assert {item["chat_id"] for item in runtime.api.sent[-2:]} == {900, -100123}
  283. await runtime.process_update(
  284. {
  285. "update_id": 22,
  286. "business_message": customer_message(
  287. message_id=2, text=None, chat_id=602, sender_id=602, photo=[{"file_id": "x"}]
  288. ),
  289. }
  290. )
  291. unsupported = await dbassistant.get_conversation_by_chat("conn-a", 602)
  292. assert unsupported["handoff_reason"] == "unsupported_message"
  293. runtime.provider = FakeProvider(error=service.AssistantProviderError("模型超时"))
  294. await runtime.process_update(
  295. {
  296. "update_id": 23,
  297. "business_message": customer_message(
  298. message_id=3, chat_id=603, sender_id=603
  299. ),
  300. }
  301. )
  302. provider_failed = await dbassistant.get_conversation_by_chat("conn-a", 603)
  303. assert provider_failed["handoff_reason"] == "provider_error"
  304. await runtime.process_update(
  305. {
  306. "update_id": 230,
  307. "business_message": customer_message(
  308. message_id=30, chat_id=605, sender_id=605
  309. ),
  310. }
  311. )
  312. quota_handoff = await dbassistant.get_conversation_by_chat("conn-a", 605)
  313. assert quota_handoff["handoff_reason"] == "account_daily_limit"
  314. sent_before = len(runtime.api.sent)
  315. await runtime.process_update(
  316. {
  317. "update_id": 24,
  318. "business_message": customer_message(
  319. message_id=4,
  320. chat_id=604,
  321. sender_id=604,
  322. sent_at=datetime.now(UTC) - timedelta(hours=25),
  323. ),
  324. }
  325. )
  326. expired = await dbassistant.get_conversation_by_chat("conn-a", 604)
  327. assert expired["handoff_reason"] == "reply_window_expired"
  328. new_sends = runtime.api.sent[sent_before:]
  329. assert len(new_sends) == 2
  330. assert all(item["business_connection_id"] == "" for item in new_sends)
  331. async def test_failed_update_goes_to_dead_letter_without_blocking(app_modules, monkeypatch):
  332. dbassistant, _service, runtime, _knowledge = await prepare_runtime(app_modules)
  333. async def fail(_message):
  334. raise RuntimeError("persistent failure")
  335. async def no_sleep(_seconds):
  336. return None
  337. runtime._handle_business_message = fail
  338. monkeypatch.setattr("wbb.services.business_assistant.asyncio.sleep", no_sleep)
  339. await runtime.process_update(
  340. {"update_id": 31, "business_message": customer_message(message_id=1)}
  341. )
  342. dead_letter = await app_modules.wbb.db.business_assistant_dead_letters.find_one(
  343. {"bot_id": "primary", "update_id": 31}
  344. )
  345. update = await app_modules.wbb.db.business_assistant_updates.find_one(
  346. {"bot_id": "primary", "update_id": 31}
  347. )
  348. assert dead_letter["error"] == "persistent failure"
  349. assert update["status"] == "done"
  350. assert await dbassistant.claim_update(31) is False
  351. async def test_connection_updates_revoke_permissions_and_persist_offset(app_modules):
  352. service = app_modules.load("wbb.services.business_assistant")
  353. dbassistant = app_modules.load("wbb.utils.dbassistant")
  354. runtime = service.BusinessAssistantRuntime(
  355. token="123456:" + "A" * 30,
  356. session=SimpleNamespace(),
  357. provider=FakeProvider(),
  358. )
  359. runtime.api = FakeBusinessApi()
  360. await runtime.process_update(
  361. {"update_id": 41, "business_connection": connection_payload(can_reply=True)}
  362. )
  363. connected = await dbassistant.get_business_connection("conn-a")
  364. assert connected["is_enabled"] is True
  365. assert connected["rights"]["can_reply"] is True
  366. await dbassistant.update_account_settings("conn-a", {"assistant_enabled": True})
  367. revoked = connection_payload(can_reply=False)
  368. revoked["is_enabled"] = False
  369. await runtime.process_update({"update_id": 42, "business_connection": revoked})
  370. disconnected = await dbassistant.get_business_connection("conn-a")
  371. assert disconnected["is_enabled"] is False
  372. assert "can_reply" not in disconnected["rights"]
  373. await runtime.process_update(
  374. {"update_id": 420, "business_message": customer_message(message_id=20)}
  375. )
  376. assert runtime.api.sent == []
  377. await dbassistant.save_update_offset(43)
  378. assert await dbassistant.load_update_offset() == 43
  379. await dbassistant.save_update_offset(44)
  380. assert await dbassistant.load_update_offset() == 44
  381. class StartupApi:
  382. def __init__(self, *, supported: bool, webhook_url: str = "") -> None:
  383. self.supported = supported
  384. self.webhook_url = webhook_url
  385. async def get_me(self):
  386. return {"username": "business_bot", "can_connect_to_business": self.supported}
  387. async def get_webhook_info(self):
  388. return {"url": self.webhook_url}
  389. async def test_runtime_blocks_webhook_conflict_and_disabled_business_mode(app_modules):
  390. service = app_modules.load("wbb.services.business_assistant")
  391. runtime = service.BusinessAssistantRuntime(
  392. token="123456:" + "A" * 30,
  393. session=SimpleNamespace(),
  394. provider=FakeProvider(),
  395. )
  396. runtime.api = StartupApi(supported=True, webhook_url="https://hooks.example.com/tg")
  397. webhook_status = await runtime.start()
  398. assert webhook_status["polling_state"] == "blocked"
  399. assert webhook_status["webhook_conflict"] is True
  400. assert runtime._poll_task is None
  401. runtime.api = StartupApi(supported=False)
  402. business_status = await runtime.start()
  403. assert business_status["polling_state"] == "blocked"
  404. assert business_status["business_mode_supported"] is False
  405. assert "BotFather" in business_status["last_error"]
  406. class FakeResponse:
  407. def __init__(self, status: int, payload: dict) -> None:
  408. self.status = status
  409. self.payload = payload
  410. async def __aenter__(self):
  411. return self
  412. async def __aexit__(self, *_args):
  413. return None
  414. async def json(self, **_kwargs):
  415. return self.payload
  416. class RetrySession:
  417. def __init__(self, responses: list[FakeResponse]) -> None:
  418. self.responses = responses
  419. self.calls = 0
  420. def post(self, *_args, **_kwargs):
  421. response = self.responses[self.calls]
  422. self.calls += 1
  423. return response
  424. async def test_telegram_api_retries_429_and_5xx(app_modules, monkeypatch):
  425. service = app_modules.load("wbb.services.business_assistant")
  426. session = RetrySession(
  427. [
  428. FakeResponse(
  429. 429,
  430. {
  431. "ok": False,
  432. "error_code": 429,
  433. "description": "Too Many Requests",
  434. "parameters": {"retry_after": 3},
  435. },
  436. ),
  437. FakeResponse(
  438. 502,
  439. {"ok": False, "error_code": 502, "description": "Bad Gateway"},
  440. ),
  441. FakeResponse(200, {"ok": True, "result": {"id": 999}}),
  442. ]
  443. )
  444. sleeps: list[int] = []
  445. async def record_sleep(seconds):
  446. sleeps.append(seconds)
  447. monkeypatch.setattr("wbb.services.business_assistant.asyncio.sleep", record_sleep)
  448. api = service.TelegramBusinessApi("123456:" + "A" * 30, session)
  449. assert await api.get_me() == {"id": 999}
  450. assert session.calls == 3
  451. assert sleeps == [3, 2]
  452. async def _login_and_change_password(client: TestClient) -> str:
  453. login = await client.post(
  454. "/api/admin/v1/auth/login",
  455. json={"username": "admin", "password": "qwe0.123456"},
  456. )
  457. login_data = (await login.json())["data"]
  458. changed = await client.put(
  459. "/api/admin/v1/auth/password",
  460. headers={"X-CSRF-Token": login_data["csrf_token"]},
  461. json={
  462. "current_password": "qwe0.123456",
  463. "new_password": "changed-pass-123",
  464. },
  465. )
  466. return (await changed.json())["data"]["csrf_token"]
  467. async def test_business_assistant_api_permission_csrf_audit_and_clear(app_modules):
  468. dbassistant = app_modules.load("wbb.utils.dbassistant")
  469. await dbassistant.upsert_business_connection(connection_payload())
  470. conversation = await dbassistant.get_or_create_conversation(
  471. "conn-a", 701, customer={"id": 701, "first_name": "待清理客户"}
  472. )
  473. await dbassistant.append_conversation_message(
  474. conversation["conversation_id"],
  475. direction="incoming",
  476. telegram_message_id=1,
  477. text="请清理",
  478. )
  479. await dbassistant.update_conversation_summary(conversation["conversation_id"], "待清理摘要")
  480. admin_api = app_modules.load("wbb.admin.api")
  481. application = admin_api.build_admin_application()
  482. await application["admin_api"].initialize()
  483. client = TestClient(TestServer(application), cookie_jar=CookieJar(unsafe=True))
  484. await client.start_server()
  485. try:
  486. csrf = await _login_and_change_password(client)
  487. app_modules.wbb.BOT_PERMISSIONS = set()
  488. denied = await client.get("/api/admin/v1/business-assistant/status")
  489. assert denied.status == 403
  490. app_modules.wbb.BOT_PERMISSIONS = {"business_assistant.manage"}
  491. no_csrf = await client.put(
  492. "/api/admin/v1/business-assistant/settings",
  493. json={"connection_id": "conn-a", "assistant_enabled": True, "confirm": True},
  494. )
  495. assert no_csrf.status == 403
  496. unconfirmed = await client.put(
  497. "/api/admin/v1/business-assistant/settings",
  498. headers={"X-CSRF-Token": csrf},
  499. json={"connection_id": "conn-a", "assistant_enabled": True},
  500. )
  501. assert unconfirmed.status == 409
  502. saved = await client.put(
  503. "/api/admin/v1/business-assistant/settings",
  504. headers={"X-CSRF-Token": csrf},
  505. json={"connection_id": "conn-a", "assistant_enabled": True, "confirm": True},
  506. )
  507. assert saved.status == 200
  508. assert (await saved.json())["data"]["assistant_enabled"] is True
  509. created = await client.post(
  510. "/api/admin/v1/business-assistant/knowledge",
  511. headers={"X-CSRF-Token": csrf},
  512. json={
  513. "connection_id": "conn-a",
  514. "question": "配送范围",
  515. "keywords": ["配送"],
  516. "answer": "仅限市区。",
  517. "confirm": True,
  518. },
  519. )
  520. assert created.status == 201
  521. listed = await client.get(
  522. "/api/admin/v1/business-assistant/knowledge?connection_id=conn-a"
  523. )
  524. assert (await listed.json())["data"]["total"] == 1
  525. cleared = await client.post(
  526. f"/api/admin/v1/business-assistant/conversations/{conversation['conversation_id']}/actions",
  527. headers={"X-CSRF-Token": csrf},
  528. json={"action": "clear", "confirm": True},
  529. )
  530. assert cleared.status == 200
  531. detail = await dbassistant.conversation_detail(conversation["conversation_id"])
  532. assert detail["summary"] == ""
  533. assert detail["messages"] == []
  534. audit = await app_modules.wbb.db.admin_audit_logs.find_one(
  535. {"action": "business_assistant.conversation.clear"}
  536. )
  537. assert audit["success"] is True
  538. finally:
  539. await client.close()