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