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

463 lines
23 KiB
Python

"""Isolated regressions for live DB reads and priority reply scheduling."""
import sqlite3
import tempfile
import threading
import time
import unittest
from pathlib import Path
from unittest import mock
from engine_b import DataEngine
from reply_database import LiveReplyDatabase
from wxwork_db import WXWorkDB, DatabaseReadError, session_fp_from_name
from wechat_bot import WeChatBot
class DatabaseFixture(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.addCleanup(self.tmp.cleanup)
self.root = Path(self.tmp.name)
self.source = self.root / "WXWork"
self.source.mkdir()
def account(self, account="100", name="客户甲"):
data = self.source / account / "Data"
data.mkdir(parents=True)
user = sqlite3.connect(data / "user.db")
user.execute("CREATE TABLE user_table(id TEXT,name TEXT,real_name TEXT,account TEXT)")
user.execute("INSERT INTO user_table VALUES('200',?,'','')", (name,))
user.commit()
user.close()
conn = sqlite3.connect(data / "message.db")
conn.execute("PRAGMA journal_mode=WAL")
conn.execute("PRAGMA wal_autocheckpoint=0")
conn.execute("CREATE TABLE message_table(sender_id TEXT,conversation_id TEXT,content_type INT,send_time INTEGER,content TEXT,server_id TEXT)")
conn.commit()
self.addCleanup(conn.close)
return conn
def insert(self, conn, stamp=100, content="新消息", sender="200", server="0", kind=2):
conn.execute("INSERT INTO message_table VALUES(?, 'M:200', ?, ?, ?, ?)", (sender, kind, stamp, content, server))
conn.commit()
def reader(self):
database = WXWorkDB(str(self.source), {}, str(self.root / "cache"))
self.addCleanup(database.close)
return database
def test_cross_thread_reads_and_wal_only_messages(self):
conn = self.account()
self.insert(conn)
database = self.reader() # Created by one thread, safely consumed by another.
result, errors = [], []
def read():
try:
result.extend(database.get_new_messages(90))
except Exception as exc:
errors.append(exc)
thread = threading.Thread(target=read)
thread.start()
thread.join(5)
self.assertFalse(thread.is_alive())
self.assertFalse(errors)
self.assertEqual([m["content"] for m in result], ["新消息"])
self.insert(conn, 100, "同秒补充")
self.assertEqual([m["content"] for m in database.get_new_messages(999)], ["同秒补充"])
def test_large_batch_and_same_second_insertion_do_not_skip(self):
conn = self.account()
conn.executemany("INSERT INTO message_table VALUES('200','M:200',2,100,?,'0')", [(str(i),) for i in range(2105)])
conn.commit()
database = self.reader()
first = database.get_new_messages(100)
second = database.get_new_messages(10000)
self.assertEqual((len(first), len(second)), (2000, 105))
self.assertEqual(len({m["dedup_key"] for m in first + second}), 2105)
self.insert(conn, 100, "最后一条")
self.assertEqual(database.get_new_messages(10000)[0]["content"], "最后一条")
self.assertEqual(database.get_new_messages(0), [])
def test_accounts_keep_independent_cursors_and_accept_late_writes(self):
first = self.account("100", "客户甲")
second = self.account("101", "客户乙")
self.insert(first, 110)
self.insert(second, 100)
database = self.reader()
self.assertEqual(len(database.get_new_messages(90)), 2)
self.insert(second, 95, "延迟落库")
self.assertEqual(database.get_new_messages(9999)[0]["content"], "延迟落库")
def test_initial_history_boundary_ms_and_new_account(self):
conn = self.account()
now = int(time.time())
self.insert(conn, now - 1000, "旧历史")
self.insert(conn, now * 1000, "毫秒时间")
database = self.reader()
result = database.get_new_messages(now - 120)
self.assertEqual([m["content"] for m in result], ["毫秒时间"])
self.assertEqual(result[0]["send_time"], now)
other = self.account("101", "客户乙")
self.insert(other, now, "新账号")
self.assertEqual(database.get_new_messages(now + 9999)[0]["content"], "新账号")
def test_context_has_sender_direction_and_media_placeholder(self):
conn = self.account()
self.insert(conn, 100, "你好")
self.insert(conn, 101, "您好", sender="100")
self.insert(conn, 102, "", kind=3)
database = self.reader()
context = database.get_conversation_context("客户甲@微信")
self.assertIn("我 1970/", context["text"])
self.assertIn("[图片]", context["text"])
self.assertFalse(context["last_message"]["is_self"])
self.assertIsNone(database.get_conversation_context("客户甲", account="999"))
self.assertIsNone(database.get_conversation_context("其他客户"))
def test_duplicate_titles_never_select_another_account(self):
for account in ("100", "101"):
self.insert(self.account(account, "同名客户"))
database = self.reader()
self.assertIsNone(database.get_conversation_context("同名客户", account="100", conv_id="M:200"))
self.assertTrue(all(m["ambiguous_name"] for m in database.get_new_messages(0)))
def test_schema_failure_does_not_advance_valid_account_cursor(self):
self.insert(self.account())
broken = self.account("101", "客户乙")
broken.execute("DROP TABLE message_table")
broken.commit()
database = self.reader()
with self.assertRaises((DatabaseReadError, sqlite3.Error)):
database.get_new_messages(0)
self.assertEqual(database._message_cursors, {})
broken.execute("CREATE TABLE message_table(sender_id TEXT,conversation_id TEXT,send_time INT,content TEXT)")
broken.commit()
self.assertEqual(database.get_new_messages(0)[0]["content"], "新消息")
def test_service_reconnect_preserves_consumed_database_rows(self):
conn = self.account()
self.insert(conn, int(time.time()))
first = self.reader()
first.get_conversation_context = mock.Mock(side_effect=DatabaseReadError("refresh failed"))
factory = mock.Mock(side_effect=[first, lambda: None])
def reopen():
if factory.call_count == 0:
factory()
return first
return WXWorkDB(str(self.source), {}, str(self.root / "cache"))
service = LiveReplyDatabase(factory=reopen, read_timeout=1, retry_seconds=0)
try:
self.assertEqual(len(service.get_new_messages(0)), 1)
self.assertIsNone(service.get_conversation_context("客户甲"))
self.assertEqual(service.get_new_messages(0), [])
self.insert(conn, int(time.time()), "重连后新消息")
self.assertEqual(service.get_new_messages(0)[0]["content"], "重连后新消息")
finally:
service.close()
service._thread.join(2)
def test_closed_database_reports_failure(self):
self.insert(self.account())
database = self.reader()
database.close()
with self.assertRaises(DatabaseReadError):
database.get_new_messages(0)
class BackgroundDatabaseTests(unittest.TestCase):
def service(self, factory, **kwargs):
service = LiveReplyDatabase(factory=factory, **kwargs)
def cleanup():
service.close()
if service._thread:
service._thread.join(2)
self.addCleanup(cleanup)
return service
def test_slow_key_setup_does_not_block_and_late_poll_result_is_retained(self):
gate = threading.Event()
database = mock.Mock()
database.get_new_messages.return_value = [{"content": "新消息"}]
threads = []
def factory():
threads.append(threading.get_ident())
gate.wait(2)
return database
service = self.service(factory, read_timeout=0.01)
started = time.monotonic()
with self.assertRaises(DatabaseReadError):
service.get_new_messages(0)
self.assertLess(time.monotonic() - started, .5)
self.assertNotEqual(threads, [threading.get_ident()])
gate.set()
result = service._pending["poll"].result(timeout=2)[1]
# Completed incremental events survive even if the UI was busy for a while.
with mock.patch("reply_database.time.monotonic", return_value=time.monotonic() + 100):
self.assertEqual(service.get_new_messages(999), result)
self.assertEqual(database.get_new_messages.call_count, 1)
def test_polling_during_retry_cooldown_does_not_extend_it(self):
database = mock.Mock()
database.get_new_messages.return_value = []
factory = mock.Mock(side_effect=[RuntimeError("等待企业微信登录"), database])
service = self.service(factory, read_timeout=1, retry_seconds=15)
with mock.patch("reply_database.time.monotonic", return_value=100):
with self.assertRaises(DatabaseReadError):
service.get_new_messages(0)
with mock.patch("reply_database.time.monotonic", return_value=110):
with self.assertRaises(DatabaseReadError):
service.get_new_messages(0)
with mock.patch("reply_database.time.monotonic", return_value=116):
self.assertEqual(service.get_new_messages(0), [])
self.assertEqual(factory.call_count, 2)
self.assertTrue(service.health_check())
def test_context_failure_returns_none_and_can_recover(self):
database = mock.Mock()
database.get_conversation_context.return_value = {"text": "你好"}
factory = mock.Mock(side_effect=[RuntimeError("暂不能解密"), database])
service = self.service(factory, read_timeout=1, retry_seconds=0)
self.assertIsNone(service.get_conversation_context("客户甲"))
self.assertEqual(service.get_conversation_context("客户甲"), {"text": "你好"})
self.assertTrue(service.health_check())
service.close()
self.assertIsNone(service.get_conversation_context("客户甲"))
class PriorityReplyTests(unittest.TestCase):
def bot(self):
bot = WeChatBot.__new__(WeChatBot)
self.fp = bytes.fromhex(session_fp_from_name("客户甲"))
bot._pending_lock = threading.RLock()
bot._pending_reply_sessions = {self.fp.hex(): {"display_name": "客户甲", "stage": "queued", "created_at": time.time()}}
bot._active_session_fp = self.fp
bot._cancelled_reply_sessions = set()
bot.identity_by_name = True
bot._persist_pending_replies = mock.Mock(return_value=True)
bot._open_chat_display_name = mock.Mock(return_value="客户甲")
bot.report_operation = mock.Mock()
bot.report_progress = mock.Mock()
bot._remember_known_outgoing_speaker = mock.Mock()
bot._forget_uncertain_tracking = mock.Mock()
bot._reply_wakeup = threading.Event()
bot._db_source = mock.Mock()
message = {"dedup_key": "new-1", "send_time": time.time(), "is_self": False}
self.context = {"text": "客户甲 2026/09/14 10:00:00\n你好", "messages": [message], "last_message": message}
bot._db_source.get_conversation_context.return_value = self.context
return bot
def test_database_success_never_touches_clipboard_or_geometry(self):
bot = self.bot()
bot._composer_geometry_valid = False
with mock.patch("wechat_bot.pyautogui.dragTo") as drag, mock.patch("wechat_bot.pyautogui.hotkey") as hotkey:
self.assertEqual(bot.extract_chat_text(), self.context["text"])
drag.assert_not_called()
hotkey.assert_not_called()
def test_current_database_context_overrides_stale_pre_copied_text(self):
bot = self.bot()
bot.store = mock.Mock()
bot.store.has_record.return_value = False
bot.extract_chat_text = mock.Mock(side_effect=AssertionError("clipboard fallback should not run"))
self.assertEqual(bot.extract_context_for(self.fp, pre_text="旧剪贴板内容"), self.context["text"])
bot.extract_chat_text.assert_not_called()
def test_actual_generation_uses_database_text_without_clipboard_or_remote_vision(self):
bot = self.bot()
bot._strict_visual_actions = True
bot.store = mock.Mock()
bot.store.has_record.return_value = False
bot._mark_reply_pending = mock.Mock(return_value=True)
bot.extract_chat_text = mock.Mock(side_effect=AssertionError("No clipboard when DB is readable"))
bot._chat_surface_signature = mock.Mock(return_value=b"surface")
bot._settled_chat_surface_signature = mock.Mock(return_value=b"surface")
bot.capture_chat_area = mock.Mock(side_effect=AssertionError("Pure DB text needs no remote image"))
bot._orchestrated_reply = mock.Mock(return_value="您好,有什么可以帮您?")
bot._stage_exchange = mock.Mock()
bot._set_task_stage = mock.Mock()
bot._log_queue_event = mock.Mock()
with mock.patch("ai_config.AI_ENABLED", True), mock.patch("ai_config.AI_USE_VISION", False), \
mock.patch("ai_config.AI_CONTEXT_ENABLED", False), mock.patch("wechat_bot.time.sleep"), \
mock.patch("registration_store.process_registration_reply", return_value=("您好,有什么可以帮您?", None)):
self.assertEqual(bot._generate_ai_reply(self.fp, chat_text="旧剪贴板"), "有什么可以帮您?")
self.assertEqual(bot._orchestrated_reply.call_args.kwargs["chat_text"], self.context["text"])
self.assertFalse(bot._orchestrated_reply.call_args.kwargs.get("force_vision", False))
bot.extract_chat_text.assert_not_called()
bot.capture_chat_area.assert_not_called()
def test_empty_database_uses_copy_fallback(self):
bot = self.bot()
bot._db_source.get_conversation_context.return_value = None
bot.store = mock.Mock()
bot.store.has_record.return_value = False
bot.extract_chat_text = mock.Mock(return_value="复制得到的新消息")
self.assertEqual(bot.extract_context_for(self.fp), "复制得到的新消息")
bot.extract_chat_text.assert_called_once_with()
def test_wrong_title_or_missing_latest_id_does_not_use_database(self):
bot = self.bot()
bot._open_chat_display_name.return_value = "客户乙"
self.assertEqual(bot._database_chat_text(self.fp), "")
bot._db_source.get_conversation_context.assert_not_called()
bot._open_chat_display_name.return_value = "客户甲"
bot._pending_reply_sessions[self.fp.hex()]["database_event"] = {"message_id": "not-in-snapshot"}
self.assertEqual(bot._database_chat_text(self.fp), "")
def test_database_direction_beats_old_visible_outgoing_bubble(self):
bot = self.bot()
text = bot._database_chat_text(self.fp)
bot._last_visible_bubble_is_outgoing = mock.Mock(return_value=True)
self.assertTrue(bot._has_pending_customer_message(text, self.fp))
self.assertFalse(bot._nothing_left_to_answer(self.fp, text))
bot._last_visible_bubble_is_outgoing.assert_not_called()
def test_followup_preserves_inflight_stage_and_wakes_worker(self):
bot = self.bot()
state = bot._pending_reply_sessions[self.fp.hex()]
state.update(stage="generating", database_event={"message_id": "old"})
bot._mark_reply_pending = mock.Mock(side_effect=AssertionError("Do not bind current UI title in detector thread"))
ok, _ = bot.enqueue_detected(self.fp.hex(), chat_text="再补充一句", database_event={"message_id": "new"})
self.assertTrue(ok)
self.assertEqual(state["stage"], "generating")
self.assertEqual(state["chat_text"], "再补充一句")
self.assertTrue(state["database_dirty"])
self.assertTrue(bot._reply_wakeup.is_set())
def test_external_label_maps_to_unique_database_task_without_changing_hashes(self):
bot = self.bot()
state = bot._pending_reply_sessions[self.fp.hex()]
state["database_event"] = {"message_id": "1", "name_unique": True}
bot._session_identity_trustworthy = mock.Mock(return_value=True)
bot._row_display_name = mock.Mock(return_value="客户甲 @微信")
visual = bot._fp_from_name("客户甲@微信")
self.assertNotEqual(visual, self.fp)
self.assertEqual(bot._session_fingerprint(None, 0), self.fp)
self.assertTrue(bot._session_fp_matches(visual, self.fp))
bot._row_display_name.return_value = "客户乙@微信"
self.assertEqual(bot._session_fingerprint(None, 0), bot._fp_from_name("客户乙@微信"))
def test_ambiguous_or_unverified_database_names_cannot_map_visual_contacts(self):
bot = self.bot()
for unique in (None, False):
bot._pending_reply_sessions[self.fp.hex()]["database_event"] = {"name_unique": unique}
self.assertIsNone(bot._database_pending_fp_for_name("客户甲@微信"))
bot._pending_reply_sessions[self.fp.hex()]["database_event"] = {"name_unique": True}
other = bot._fp_from_name("客户甲@微信")
bot._pending_reply_sessions[other.hex()] = {"display_name": "客户甲@微信", "database_event": {"name_unique": True}}
self.assertIsNone(bot._database_pending_fp_for_name("客户甲@微信"))
def test_database_event_merges_into_existing_external_visual_task(self):
bot = self.bot()
visual = bot._fp_from_name("客户甲@微信")
bot._pending_reply_sessions = {visual.hex(): {"display_name": "客户甲@微信", "stage": "generating"}}
self.assertTrue(bot.enqueue_detected(self.fp.hex(), display_name="客户甲", chat_text="补充消息",
database_event={"message_id": "2", "name_unique": True})[0])
self.assertEqual(set(bot._pending_reply_sessions), {visual.hex()})
self.assertTrue(bot._pending_reply_sessions[visual.hex()]["database_dirty"])
def test_new_task_does_not_bind_foreground_identity(self):
bot = self.bot()
bot._pending_reply_sessions.clear()
def mark(fp, **kwargs):
self.assertIs(kwargs["bind_identity"], False)
bot._pending_reply_sessions[fp.hex()] = {"display_name": "客户甲", "stage": "queued"}
return True
bot._mark_reply_pending = mock.Mock(side_effect=mark)
self.assertTrue(bot.enqueue_detected(self.fp.hex(), chat_text="你好")[0])
bot._mark_reply_pending.assert_called_once()
def test_followup_invalidates_reply_generated_for_old_message(self):
bot = self.bot()
state = bot._pending_reply_sessions[self.fp.hex()]
state["database_event"] = {"message_id": "old"}
def generate(*args, **kwargs):
state.update(staged_reply_text="旧回复", database_event={"message_id": "new"}, database_dirty=True)
return "旧回复"
bot._generate_ai_reply_impl = mock.Mock(side_effect=generate)
self.assertIsNone(bot._generate_ai_reply(self.fp))
self.assertNotIn("staged_reply_text", state)
self.assertEqual(state["database_consumed_id"], "new")
def test_uncertain_send_is_not_reset_by_followup(self):
bot = self.bot()
state = bot._pending_reply_sessions[self.fp.hex()]
state.update(send_state="uncertain", staged_reply_text="等待回执", database_event={"message_id": "new"}, database_dirty=True)
bot._consume_database_update(self.fp)
self.assertEqual(state["send_state"], "uncertain")
self.assertEqual(state["staged_reply_text"], "等待回执")
self.assertTrue(state["database_dirty"])
def test_new_message_arriving_before_send_rejects_stale_reply(self):
bot = self.bot()
bot._pending_reply_sessions[self.fp.hex()].update(generation_database_id="old", database_event={"message_id": "new"})
bot._deny_send = mock.Mock(return_value=False)
with mock.patch("wechat_bot.pyautogui.hotkey") as hotkey:
self.assertFalse(bot.send_reply("旧回复", expected_fp=self.fp))
hotkey.assert_not_called()
bot._deny_send.assert_called_once()
def test_failed_active_task_yields_to_new_customer(self):
bot = self.bot()
other = bytes.fromhex(session_fp_from_name("客户乙"))
state = bot._pending_reply_sessions[self.fp.hex()]
state["resume_failures"] = 1
bot._pending_reply_sessions[other.hex()] = {"created_at": time.time(), "database_event": {"message_id": "new"}}
self.assertEqual(bot._pending_queue_order(bot._pending_reply_sessions, self.fp)[0][0], other.hex())
self.assertFalse(bot._hold_for_unfinished_active_session())
class DatabaseDetectionTests(unittest.TestCase):
def engine(self):
bot = mock.Mock()
bot.has_active_pending.return_value = True
bot.enqueue_detected.return_value = (True, "已投递")
database = mock.Mock()
engine = DataEngine(bot=bot, db_source=database)
self.bot, self.database = bot, database
self.fp = session_fp_from_name("客户甲")
return engine
def message(self, ident=1, own=False):
return {"fp_hex": self.fp, "account": "100", "conv_id": "M:200", "display_name": "客户甲", "content": "新消息", "send_time": time.time(), "dedup_key": str(ident), "rowid": ident, "is_self": own}
def test_active_task_receives_followup_instead_of_only_merging_tag(self):
engine = self.engine()
self.database.get_new_messages.side_effect = [[self.message(1)], [self.message(2)]]
engine._poll_db_once()
engine._poll_db_once()
self.assertEqual(self.bot.enqueue_detected.call_count, 2)
self.assertEqual(self.bot.enqueue_detected.call_args.kwargs["database_event"]["message_id"], "2")
self.bot.merge_detected_by.assert_not_called()
def test_failed_enqueue_is_retried_without_new_database_row(self):
engine = self.engine()
self.bot.enqueue_detected.side_effect = [(False, "busy"), (True, "ok")]
self.database.get_new_messages.side_effect = [[self.message()], []]
engine._poll_db_once()
self.assertEqual(len(engine._pending_db_events), 1)
engine._poll_db_once()
self.assertEqual(engine._pending_db_events, {})
self.assertEqual(engine.enqueued_count, 1)
def test_last_outgoing_is_not_queued_and_failure_changes_health(self):
engine = self.engine()
self.database.get_new_messages.side_effect = [[self.message(1), self.message(2, own=True)], DatabaseReadError("unreadable"), []]
engine._poll_db_once()
self.bot.enqueue_detected.assert_not_called()
engine._poll_db_once()
self.assertFalse(engine.db_active)
engine._poll_db_once()
self.assertTrue(engine.db_active)
def test_duplicate_message_and_ambiguous_name_are_not_queued(self):
engine = self.engine()
msg = self.message()
self.database.get_new_messages.side_effect = [[msg], [msg], [{**self.message(2), "ambiguous_name": True}]]
for _ in range(3):
engine._poll_db_once()
self.assertEqual(self.bot.enqueue_detected.call_count, 1)
if __name__ == "__main__":
unittest.main()