112 lines
3.6 KiB
Python
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()
|