迁移目录
This commit is contained in:
@@ -0,0 +1,111 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user