463 lines
23 KiB
Python
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()
|