"""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()