Files
sbnews/server/public/dongqiudi-crawler/tests/test_kb_worker.py
T

112 lines
3.6 KiB
Python

import unittest
from pathlib import Path
import sys
import types
APP_DIR = Path(__file__).resolve().parents[1]
if str(APP_DIR) not in sys.path:
sys.path.insert(0, str(APP_DIR))
fake_aiomysql = types.ModuleType("aiomysql")
fake_aiomysql.DictCursor = object
fake_aiomysql.create_pool = None
sys.modules.setdefault("aiomysql", fake_aiomysql)
fake_aiohttp = types.ModuleType("aiohttp")
fake_aiohttp.ClientSession = object
fake_aiohttp.ClientTimeout = object
sys.modules.setdefault("aiohttp", fake_aiohttp)
fake_redis = types.ModuleType("redis")
fake_redis_asyncio = types.ModuleType("redis.asyncio")
fake_redis_asyncio.Redis = object
fake_redis.asyncio = fake_redis_asyncio
sys.modules.setdefault("redis", fake_redis)
sys.modules.setdefault("redis.asyncio", fake_redis_asyncio)
fake_loguru = types.ModuleType("loguru")
fake_loguru.logger = types.SimpleNamespace(info=lambda *args, **kwargs: None, warning=lambda *args, **kwargs: None, error=lambda *args, **kwargs: None)
sys.modules.setdefault("loguru", fake_loguru)
fake_src = types.ModuleType("src")
fake_core = types.ModuleType("src.core")
fake_config = types.ModuleType("src.core.config")
fake_config.get_config = lambda: None
fake_config.load_config = lambda: None
fake_core.config = fake_config
fake_src.core = fake_core
sys.modules.setdefault("src", fake_src)
sys.modules.setdefault("src.core", fake_core)
sys.modules.setdefault("src.core.config", fake_config)
from scripts.kb_worker import DEAD_STREAM_KEY, GROUP_NAME, STREAM_KEY, KbWorker # noqa: E402
class FakeRedis:
def __init__(self):
self.acked = []
self.added = []
async def xack(self, stream, group, message_id):
self.acked.append((stream, group, message_id))
async def xadd(self, stream, payload):
self.added.append((stream, payload))
return "2-0"
class FakeDb:
def __init__(self, result=None, exc=None):
self.result = result if result is not None else {"success": True, "document_id": 1, "chunk_count": 1}
self.exc = exc
self.calls = []
async def upsert_document(self, domain, subtype, source_id, embedder):
self.calls.append((domain, subtype, source_id))
if self.exc:
raise self.exc
return self.result
class FakeEmbedder:
pass
class KbWorkerQueueTest(unittest.IsolatedAsyncioTestCase):
def make_worker(self, db):
worker = object.__new__(KbWorker)
worker.redis = FakeRedis()
worker.db = db
worker.embedder = FakeEmbedder()
worker.max_retries = 1
return worker
async def test_successful_message_is_acked(self):
worker = self.make_worker(FakeDb())
await worker.handle_message("1-0", {"domain": "article", "subtype": "article", "source_id": "123"})
self.assertEqual(worker.db.calls, [("article", "article", 123)])
self.assertEqual(worker.redis.acked, [(STREAM_KEY, GROUP_NAME, "1-0")])
self.assertEqual(worker.redis.added, [])
async def test_message_goes_to_dead_stream_after_retry_limit(self):
worker = self.make_worker(FakeDb(exc=RuntimeError("boom")))
await worker.handle_message(
"1-0",
{"domain": "article", "subtype": "article", "source_id": "123", "attempts": "1"},
)
self.assertEqual(worker.redis.acked, [(STREAM_KEY, GROUP_NAME, "1-0")])
self.assertEqual(len(worker.redis.added), 1)
stream, payload = worker.redis.added[0]
self.assertEqual(stream, DEAD_STREAM_KEY)
self.assertEqual(payload["attempts"], "2")
self.assertIn("boom", payload["last_error"])
if __name__ == "__main__":
unittest.main()