468 lines
17 KiB
Python
468 lines
17 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import os
|
|
import sys
|
|
import unittest
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock, Mock
|
|
|
|
|
|
BACKEND_DIR = Path(__file__).resolve().parents[1]
|
|
os.environ["KEFU_DB_TYPE"] = "sqlite"
|
|
os.environ["KEFU_DATABASE_URL"] = ""
|
|
os.environ["KEFU_DB_PATH"] = str(BACKEND_DIR / "kefu.db")
|
|
if str(BACKEND_DIR) not in sys.path:
|
|
sys.path.insert(0, str(BACKEND_DIR))
|
|
|
|
from auth.system_settings import SystemSettingsData, set_cached_settings
|
|
from rpa_engine.douyin_im.service import DouyinImService
|
|
from rpa_engine.playwright_worker import DouyinWorker
|
|
|
|
|
|
class _RecordingQueue:
|
|
def __init__(self) -> None:
|
|
self.jobs: list[
|
|
tuple[float, object, str, dict, frozenset[str], bool]
|
|
] = []
|
|
|
|
async def enqueue(
|
|
self,
|
|
delay_seconds,
|
|
callback,
|
|
description="",
|
|
details=None,
|
|
merge_key="",
|
|
merge_keys=None,
|
|
immediate_if_idle=False,
|
|
) -> int:
|
|
keys = merge_keys if merge_keys is not None else [merge_key]
|
|
normalized_keys = frozenset(str(key) for key in keys if str(key or "").strip())
|
|
self.jobs.append(
|
|
(
|
|
delay_seconds,
|
|
callback,
|
|
description,
|
|
dict(details or {}),
|
|
normalized_keys,
|
|
bool(immediate_if_idle),
|
|
)
|
|
)
|
|
return len(self.jobs)
|
|
|
|
async def merge_pending(self, merge_key, details_merger):
|
|
incoming_keys = (
|
|
frozenset([merge_key])
|
|
if isinstance(merge_key, str)
|
|
else frozenset(merge_key)
|
|
)
|
|
for index, job in enumerate(self.jobs):
|
|
existing_conversations = {
|
|
key for key in job[4] if key.startswith("conv:")
|
|
}
|
|
incoming_conversations = {
|
|
key for key in incoming_keys if key.startswith("conv:")
|
|
}
|
|
same_conversation = bool(
|
|
existing_conversations & incoming_conversations
|
|
)
|
|
different_explicit_conversations = bool(
|
|
existing_conversations
|
|
and incoming_conversations
|
|
and not same_conversation
|
|
)
|
|
same_peer = bool(
|
|
{key for key in job[4] if key.startswith("peer:")}
|
|
& {key for key in incoming_keys if key.startswith("peer:")}
|
|
)
|
|
if not same_conversation and (different_explicit_conversations or not same_peer):
|
|
continue
|
|
merged_details = details_merger(dict(job[3]))
|
|
self.jobs[index] = (
|
|
job[0],
|
|
job[1],
|
|
job[2],
|
|
merged_details,
|
|
frozenset(job[4] | incoming_keys),
|
|
job[5],
|
|
)
|
|
return {
|
|
"status": "merged",
|
|
"job_id": f"job-{index + 1}",
|
|
"position": index + 1,
|
|
"message_count": merged_details.get("message_count", 1),
|
|
}
|
|
return {"status": "not_found"}
|
|
|
|
|
|
def _build_service(delay_seconds: int = 60):
|
|
match_reply = AsyncMock(return_value=["自动回复"])
|
|
log_fn = AsyncMock()
|
|
service = DouyinImService(
|
|
session=SimpleNamespace(my_uid="999"),
|
|
match_reply=match_reply,
|
|
log_fn=log_fn,
|
|
account_id=1,
|
|
)
|
|
service._running = True
|
|
service._reply_queue = _RecordingQueue()
|
|
service._resolve_peer_profile = AsyncMock(return_value=("张三", "", "123"))
|
|
service._resolve_cooldown_seconds = AsyncMock(return_value=0)
|
|
service._resolve_reply_delay_seconds = AsyncMock(return_value=delay_seconds)
|
|
return service, match_reply, log_fn
|
|
|
|
|
|
class ReplyQueueIntegrationTests(unittest.IsolatedAsyncioTestCase):
|
|
async def test_kick_response_takes_account_offline_immediately(self):
|
|
callback = AsyncMock()
|
|
service, _, _ = _build_service()
|
|
service.on_session_invalid = callback
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._note_session_invalid(
|
|
"抖音安全网关返回 decision=KICK,当前登录态已失效"
|
|
)
|
|
|
|
self.assertFalse(service._running)
|
|
self.assertTrue(service._session_invalid_fired)
|
|
callback.assert_awaited_once()
|
|
self.assertIn("decision=KICK", callback.await_args.args[0])
|
|
|
|
async def test_invalid_request_still_requires_two_consecutive_failures(self):
|
|
callback = AsyncMock()
|
|
service, _, _ = _build_service()
|
|
service.on_session_invalid = callback
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._note_session_invalid("INVALID_REQUEST")
|
|
self.assertTrue(service._running)
|
|
callback.assert_not_awaited()
|
|
await service._note_session_invalid("INVALID_REQUEST")
|
|
|
|
self.assertFalse(service._running)
|
|
callback.assert_awaited_once()
|
|
|
|
async def test_same_message_from_ws_and_poll_is_queued_once(self):
|
|
service, match_reply, _ = _build_service()
|
|
message = {
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
"content": "你好",
|
|
"server_message_id": "mid-1",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await asyncio.gather(
|
|
service._handle_incoming(dict(message)),
|
|
service._handle_incoming(dict(message)),
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 1)
|
|
self.assertEqual(match_reply.await_count, 1)
|
|
self.assertEqual(service._reply_queue.jobs[0][0], 60)
|
|
self.assertTrue(service._reply_queue.jobs[0][5])
|
|
details = service._reply_queue.jobs[0][3]
|
|
self.assertEqual(details["sender_name"], "张三")
|
|
self.assertEqual(details["conversation_id"], "conv-1")
|
|
self.assertEqual(details["incoming_content"], "你好")
|
|
self.assertEqual(details["replies"], ["自动回复"])
|
|
|
|
async def test_same_conversation_with_different_message_ids_merges(self):
|
|
service, match_reply, _ = _build_service()
|
|
base = {
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
"content": "你好",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming({**base, "server_message_id": "mid-1"})
|
|
await service._handle_incoming({**base, "server_message_id": "mid-2"})
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 1)
|
|
self.assertEqual(match_reply.await_count, 1)
|
|
details = service._reply_queue.jobs[0][3]
|
|
self.assertEqual(details["incoming_contents"], ["你好", "你好"])
|
|
self.assertEqual(details["message_count"], 2)
|
|
self.assertEqual(details["replies"], ["自动回复"])
|
|
|
|
async def test_pending_conversation_merges_before_cooldown_filter(self):
|
|
service, match_reply, _ = _build_service()
|
|
service._resolve_cooldown_seconds = AsyncMock(return_value=120)
|
|
base = {
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming(
|
|
{**base, "content": "first", "server_message_id": "mid-1"}
|
|
)
|
|
await service._handle_incoming(
|
|
{
|
|
**base,
|
|
"conversation_id": "conv-1",
|
|
"content": "second",
|
|
"server_message_id": "mid-2",
|
|
}
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 1)
|
|
self.assertEqual(match_reply.await_count, 1)
|
|
details = service._reply_queue.jobs[0][3]
|
|
self.assertEqual(details["incoming_contents"], ["first", "second"])
|
|
self.assertEqual(details["message_count"], 2)
|
|
|
|
async def test_http_previews_from_different_conversations_do_not_collide(self):
|
|
service, _, _ = _build_service()
|
|
base = {
|
|
"sender_uid": "",
|
|
"sender_name": "同名用户",
|
|
"content": "相同内容",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming({**base, "conversation_id": "conv-a"})
|
|
await service._handle_incoming({**base, "conversation_id": "conv-b"})
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 2)
|
|
|
|
async def test_missing_conversation_id_uses_peer_uid_for_merge(self):
|
|
service, match_reply, _ = _build_service()
|
|
base = {
|
|
"conversation_id": "",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming(
|
|
{**base, "content": "first", "server_message_id": "mid-1"}
|
|
)
|
|
await service._handle_incoming(
|
|
{
|
|
**base,
|
|
"conversation_id": "conv-1",
|
|
"content": "second",
|
|
"server_message_id": "mid-2",
|
|
}
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 1)
|
|
self.assertEqual(match_reply.await_count, 1)
|
|
self.assertEqual(
|
|
service._reply_queue.jobs[0][3]["incoming_contents"],
|
|
["first", "second"],
|
|
)
|
|
|
|
async def test_peer_alias_added_later_merges_with_conversation_only_job(self):
|
|
service, match_reply, _ = _build_service()
|
|
|
|
async def resolve_profile(conv_id, sender_uid, sender, sender_avatar):
|
|
return "张三", "", str(sender_uid or "")
|
|
|
|
service._resolve_peer_profile = AsyncMock(side_effect=resolve_profile)
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming(
|
|
{
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "",
|
|
"sender_name": "张三",
|
|
"content": "first",
|
|
"server_message_id": "mid-1",
|
|
}
|
|
)
|
|
await service._handle_incoming(
|
|
{
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
"content": "second",
|
|
"server_message_id": "mid-2",
|
|
}
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 1)
|
|
self.assertEqual(match_reply.await_count, 1)
|
|
self.assertEqual(
|
|
service._reply_queue.jobs[0][3]["incoming_contents"],
|
|
["first", "second"],
|
|
)
|
|
|
|
async def test_missing_conversation_and_peer_ids_do_not_merge_by_name(self):
|
|
service, match_reply, _ = _build_service()
|
|
service._resolve_peer_profile = AsyncMock(return_value=("同名用户", "", ""))
|
|
base = {
|
|
"conversation_id": "",
|
|
"sender_uid": "",
|
|
"sender_name": "同名用户",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming(
|
|
{**base, "content": "first", "server_message_id": "mid-1"}
|
|
)
|
|
await service._handle_incoming(
|
|
{**base, "content": "second", "server_message_id": "mid-2"}
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 2)
|
|
self.assertEqual(match_reply.await_count, 2)
|
|
|
|
async def test_slow_first_message_keeps_arrival_order(self):
|
|
service, _, _ = _build_service()
|
|
|
|
async def resolve_profile(conv_id, sender_uid, sender, sender_avatar):
|
|
if conv_id == "conv-first":
|
|
await asyncio.sleep(0.03)
|
|
return "先到", "", "101"
|
|
return "后到", "", "102"
|
|
|
|
service._resolve_peer_profile = AsyncMock(side_effect=resolve_profile)
|
|
first = {
|
|
"conversation_id": "conv-first",
|
|
"sender_uid": "101",
|
|
"sender_name": "先到",
|
|
"content": "第一条",
|
|
"server_message_id": "mid-first",
|
|
}
|
|
second = {
|
|
"conversation_id": "conv-second",
|
|
"sender_uid": "102",
|
|
"sender_name": "后到",
|
|
"content": "第二条",
|
|
"server_message_id": "mid-second",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
first_task = asyncio.create_task(service._handle_incoming(first))
|
|
await asyncio.sleep(0)
|
|
second_task = asyncio.create_task(service._handle_incoming(second))
|
|
await asyncio.gather(first_task, second_task)
|
|
|
|
descriptions = [job[2] for job in service._reply_queue.jobs]
|
|
self.assertEqual(descriptions, ["回复 先到", "回复 后到"])
|
|
|
|
async def test_zero_effective_delay_skips_queue_and_sends_immediately(self):
|
|
service, _, _ = _build_service(delay_seconds=0)
|
|
service._send_auto_reply = AsyncMock()
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming(
|
|
{
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
"content": "你好",
|
|
"server_message_id": "mid-1",
|
|
}
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 0)
|
|
service._send_auto_reply.assert_awaited_once()
|
|
|
|
async def test_zero_delay_keeps_separate_immediate_replies(self):
|
|
service, _, _ = _build_service(delay_seconds=0)
|
|
service._send_auto_reply = AsyncMock()
|
|
base = {
|
|
"conversation_id": "conv-1",
|
|
"sender_uid": "123",
|
|
"sender_name": "张三",
|
|
}
|
|
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._handle_incoming(
|
|
{**base, "content": "first", "server_message_id": "mid-1"}
|
|
)
|
|
await service._handle_incoming(
|
|
{**base, "content": "second", "server_message_id": "mid-2"}
|
|
)
|
|
|
|
self.assertEqual(len(service._reply_queue.jobs), 0)
|
|
self.assertEqual(service._send_auto_reply.await_count, 2)
|
|
|
|
async def test_session_stop_during_first_payload_cancels_remaining_payloads(self):
|
|
service, _, log_fn = _build_service()
|
|
|
|
async def fail_and_stop(*args, **kwargs):
|
|
service.last_error = "INVALID_REQUEST"
|
|
service._running = False
|
|
return False, None
|
|
|
|
service._send_text = AsyncMock(side_effect=fail_and_stop)
|
|
with unittest.mock.patch(
|
|
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
|
):
|
|
await service._send_auto_reply(
|
|
sender="张三",
|
|
content="你好",
|
|
conv_id="conv-1",
|
|
replies=["第一段", "第二段"],
|
|
peer_key="123",
|
|
cooldown=0,
|
|
log_kwargs={
|
|
"sender_name": "张三",
|
|
"sender_id": "123",
|
|
"sender_avatar": None,
|
|
"message": "你好",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(service._send_text.await_count, 1)
|
|
log_fn.assert_awaited_once()
|
|
|
|
|
|
class ReplyDelayResolutionTests(unittest.IsolatedAsyncioTestCase):
|
|
async def asyncSetUp(self):
|
|
self.worker = DouyinWorker(account_id=1)
|
|
|
|
async def test_account_value_overrides_system_default(self):
|
|
set_cached_settings(SystemSettingsData(auto_reply_delay_seconds=60))
|
|
self.worker.get_reply_delay = AsyncMock(return_value=15)
|
|
|
|
self.assertEqual(await self.worker.resolve_reply_delay_seconds(), 15)
|
|
|
|
async def test_unset_account_uses_system_default(self):
|
|
set_cached_settings(SystemSettingsData(auto_reply_delay_seconds=60))
|
|
self.worker.get_reply_delay = AsyncMock(return_value=None)
|
|
|
|
self.assertEqual(await self.worker.resolve_reply_delay_seconds(), 60)
|
|
|
|
async def test_both_unset_skip_queue_rule(self):
|
|
set_cached_settings(SystemSettingsData(auto_reply_delay_seconds=0))
|
|
self.worker.get_reply_delay = AsyncMock(return_value=None)
|
|
|
|
self.assertEqual(await self.worker.resolve_reply_delay_seconds(), 0)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|