迁移目录
This commit is contained in:
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
import unittest
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
if str(ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
from ai_comment_dispatch import (
|
||||
AiCommentDispatcher,
|
||||
AiCommentSettings,
|
||||
CommentCandidate,
|
||||
DispatchTask,
|
||||
VirtualUser,
|
||||
extract_post_text,
|
||||
merge_candidates,
|
||||
select_balanced_candidates,
|
||||
)
|
||||
|
||||
|
||||
class FakeRepo:
|
||||
def __init__(self, article_candidates=None, post_candidates=None, daily_success=0, users=None):
|
||||
self.article_candidates = article_candidates or []
|
||||
self.post_candidates = post_candidates or []
|
||||
self.daily_success = daily_success
|
||||
self.users = users or []
|
||||
self.created_tasks = []
|
||||
self.saved_comments = []
|
||||
self.failed_tasks = []
|
||||
self.ensure_schema_called = False
|
||||
self.ensure_users_target = None
|
||||
|
||||
def ensure_schema(self):
|
||||
self.ensure_schema_called = True
|
||||
|
||||
def ensure_virtual_users(self, target_count):
|
||||
self.ensure_users_target = target_count
|
||||
return list(self.users)
|
||||
|
||||
def count_daily_success(self, day_start, day_end):
|
||||
return self.daily_success
|
||||
|
||||
def load_article_candidates(self, settings):
|
||||
return list(self.article_candidates)
|
||||
|
||||
def load_post_candidates(self, settings):
|
||||
return list(self.post_candidates)
|
||||
|
||||
def create_task(self, candidate, virtual_user, persona_key):
|
||||
task = DispatchTask(
|
||||
task_id=len(self.created_tasks) + 1,
|
||||
target_type=candidate.target_type,
|
||||
target_id=candidate.target_id,
|
||||
virtual_user_id=virtual_user.user_id,
|
||||
persona_key=persona_key,
|
||||
selection_sources=tuple(candidate.selection_sources),
|
||||
)
|
||||
self.created_tasks.append(task)
|
||||
return task
|
||||
|
||||
def mark_task_success(self, task, comment_content):
|
||||
self.saved_comments.append((task, comment_content))
|
||||
|
||||
def mark_task_failed(self, task, message):
|
||||
self.failed_tasks.append((task, message))
|
||||
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, outputs=None):
|
||||
self.outputs = outputs or {}
|
||||
self.calls = []
|
||||
|
||||
def generate_comment(self, candidate, virtual_user, persona_key, settings):
|
||||
self.calls.append((candidate, virtual_user, persona_key))
|
||||
return self.outputs.get(
|
||||
(candidate.target_type, candidate.target_id),
|
||||
{"success": True, "content": f"{candidate.target_type}-{candidate.target_id}-comment"},
|
||||
)
|
||||
|
||||
|
||||
class AiCommentDispatchTest(unittest.TestCase):
|
||||
def test_merge_candidates_merges_selection_sources_by_target(self):
|
||||
merged = merge_candidates(
|
||||
[
|
||||
CommentCandidate("article", 10, "title-a", "body-a", ("cate:1",), created_at=100, rank_value=8),
|
||||
CommentCandidate("article", 10, "title-a", "body-a", ("top",), created_at=100, rank_value=8),
|
||||
CommentCandidate("article", 11, "title-b", "body-b", ("hot",), created_at=90, rank_value=5),
|
||||
]
|
||||
)
|
||||
|
||||
self.assertEqual(len(merged), 2)
|
||||
self.assertEqual(merged[0].target_id, 10)
|
||||
self.assertEqual(set(merged[0].selection_sources), {"cate:1", "top"})
|
||||
|
||||
def test_select_balanced_candidates_prefers_25_25_then_backfills(self):
|
||||
articles = [
|
||||
CommentCandidate("article", index, f"title-{index}", "body", (f"cate:{index}",), created_at=200 - index, rank_value=index)
|
||||
for index in range(1, 9)
|
||||
]
|
||||
posts = [
|
||||
CommentCandidate("post", index, "", f"post-{index}", (f"tag:{index}",), created_at=100 - index, rank_value=index)
|
||||
for index in range(1, 3)
|
||||
]
|
||||
|
||||
selected = select_balanced_candidates(articles, posts, total_limit=6)
|
||||
|
||||
self.assertEqual(len(selected), 6)
|
||||
self.assertEqual(sum(1 for item in selected if item.target_type == "post"), 2)
|
||||
self.assertEqual(sum(1 for item in selected if item.target_type == "article"), 4)
|
||||
self.assertEqual([item.target_id for item in selected[:3]], [1, 2, 3])
|
||||
|
||||
def test_extract_post_text_prefers_image_only_analysis_content(self):
|
||||
content = extract_post_text(
|
||||
"",
|
||||
{
|
||||
"lottery_analysis_source": "image_only",
|
||||
"lottery_analysis_content": "杀码思路和冷热分析",
|
||||
},
|
||||
)
|
||||
self.assertEqual(content, "杀码思路和冷热分析")
|
||||
|
||||
def test_dispatcher_stops_when_daily_cap_reached(self):
|
||||
repo = FakeRepo(daily_success=500, users=[VirtualUser(1, "ai_comment_0001", "AI评论员01", "data_analyst")])
|
||||
client = FakeClient()
|
||||
dispatcher = AiCommentDispatcher(repo, client)
|
||||
settings = AiCommentSettings(per_run=10, daily_cap=500, user_pool_size=100)
|
||||
|
||||
summary = dispatcher.run(settings)
|
||||
|
||||
self.assertTrue(repo.ensure_schema_called)
|
||||
self.assertEqual(repo.ensure_users_target, 100)
|
||||
self.assertEqual(summary["status"], "skipped")
|
||||
self.assertEqual(summary["reason"], "daily_cap_reached")
|
||||
self.assertEqual(len(repo.created_tasks), 0)
|
||||
self.assertEqual(len(client.calls), 0)
|
||||
|
||||
def test_dispatcher_creates_tasks_and_marks_success_or_failure(self):
|
||||
article_candidates = [
|
||||
CommentCandidate("article", 101, "article-101", "body", ("cate:10",), created_at=10, rank_value=3),
|
||||
CommentCandidate("article", 102, "article-102", "body", ("hot",), created_at=9, rank_value=2),
|
||||
]
|
||||
post_candidates = [
|
||||
CommentCandidate("post", 201, "", "post-201", ("tag:14",), created_at=8, rank_value=5),
|
||||
CommentCandidate("post", 202, "", "post-202", ("tag:15",), created_at=7, rank_value=4),
|
||||
]
|
||||
users = [
|
||||
VirtualUser(1, "ai_comment_0001", "AI评论员01", "data_analyst"),
|
||||
VirtualUser(2, "ai_comment_0002", "AI评论员02", "old_fan"),
|
||||
VirtualUser(3, "ai_comment_0003", "AI评论员03", "coach_view"),
|
||||
VirtualUser(4, "ai_comment_0004", "AI评论员04", "night_shift_editor"),
|
||||
]
|
||||
repo = FakeRepo(article_candidates=article_candidates, post_candidates=post_candidates, users=users)
|
||||
client = FakeClient(
|
||||
outputs={
|
||||
("article", 101): {"success": True, "content": "文章评论A"},
|
||||
("article", 102): {"success": False, "error": "上游超时"},
|
||||
("post", 201): {"success": True, "content": "帖子评论B"},
|
||||
("post", 202): {"success": True, "content": "帖子评论C"},
|
||||
}
|
||||
)
|
||||
dispatcher = AiCommentDispatcher(repo, client)
|
||||
settings = AiCommentSettings(per_run=4, daily_cap=500, user_pool_size=4, concurrency=2)
|
||||
|
||||
summary = dispatcher.run(settings)
|
||||
|
||||
self.assertEqual(summary["status"], "success")
|
||||
self.assertEqual(summary["selected_count"], 4)
|
||||
self.assertEqual(summary["success_count"], 4)
|
||||
self.assertEqual(summary["failed_count"], 0)
|
||||
self.assertEqual(len(repo.created_tasks), 4)
|
||||
self.assertEqual(len(repo.saved_comments), 4)
|
||||
self.assertEqual(len(repo.failed_tasks), 0)
|
||||
self.assertIn("后续", repo.saved_comments[1][1])
|
||||
self.assertEqual(len({task.virtual_user_id for task in repo.created_tasks}), 4)
|
||||
|
||||
def test_settings_reads_new_config_values(self):
|
||||
settings = AiCommentSettings.from_config_map(
|
||||
{
|
||||
"auto_comment_enabled": "1",
|
||||
"auto_comment_user_pool_size": "100",
|
||||
"auto_comment_per_run": "50",
|
||||
"auto_comment_daily_cap": "500",
|
||||
"auto_comment_concurrency": "4",
|
||||
"auto_comment_request_timeout": "180",
|
||||
"auto_comment_article_prompt": "article prompt",
|
||||
"auto_comment_post_prompt": "post prompt",
|
||||
"openai_api_key": "sk-demo",
|
||||
"openai_base_url": "https://example.com",
|
||||
"openai_model": "gpt-5.5",
|
||||
"openai_wire_api": "responses",
|
||||
"openai_provider_label": "OpenAI",
|
||||
"openai_reasoning_effort": "high",
|
||||
"openai_disable_response_storage": "1",
|
||||
"max_tokens": "300",
|
||||
"temperature": "0.65",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertTrue(settings.enabled)
|
||||
self.assertEqual(settings.user_pool_size, 100)
|
||||
self.assertEqual(settings.per_run, 50)
|
||||
self.assertEqual(settings.daily_cap, 500)
|
||||
self.assertEqual(settings.concurrency, 4)
|
||||
self.assertEqual(settings.request_timeout, 180)
|
||||
self.assertEqual(settings.article_prompt, "article prompt")
|
||||
self.assertEqual(settings.post_prompt, "post prompt")
|
||||
self.assertEqual(settings.openai_api_key, "sk-demo")
|
||||
self.assertEqual(settings.max_tokens, 300)
|
||||
self.assertAlmostEqual(settings.temperature, 0.65)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,203 @@
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
import importlib
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
if str(ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
|
||||
class FakeStore:
|
||||
def __init__(self):
|
||||
self.created = []
|
||||
self.progress = []
|
||||
self.finished = []
|
||||
self.next_id = 100
|
||||
|
||||
def create_running(self, task):
|
||||
self.created.append(task)
|
||||
self.next_id += 1
|
||||
return self.next_id
|
||||
|
||||
def update_progress(self, log_id, output, elapsed):
|
||||
self.progress.append((log_id, output, elapsed))
|
||||
|
||||
def finish(self, log_id, status, output, elapsed, error_message=""):
|
||||
self.finished.append((log_id, status, output, elapsed, error_message))
|
||||
|
||||
|
||||
class DockerTaskRunnerTest(unittest.TestCase):
|
||||
def test_config_uses_dqd_environment_overrides(self):
|
||||
sys.modules.pop("src.core.config", None)
|
||||
src_core = importlib.import_module("src.core")
|
||||
if hasattr(src_core, "config"):
|
||||
delattr(src_core, "config")
|
||||
config_module = importlib.import_module("src.core.config")
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
settings = Path(tmp) / "settings.yaml"
|
||||
settings.write_text(
|
||||
"\n".join(
|
||||
[
|
||||
"database:",
|
||||
' host: "yaml-db"',
|
||||
" port: 3306",
|
||||
' database: "yaml_name"',
|
||||
' username: "yaml_user"',
|
||||
' password: "yaml_password"',
|
||||
' prefix: "la_"',
|
||||
"redis:",
|
||||
' host: "yaml-redis"',
|
||||
" port: 6379",
|
||||
" db: 0",
|
||||
]
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
env = {
|
||||
"DQD_DB_HOST": "docker-db",
|
||||
"DQD_DB_PORT": "3307",
|
||||
"DQD_DB_NAME": "docker_name",
|
||||
"DQD_DB_USER": "docker_user",
|
||||
"DQD_DB_PASSWORD": "docker_password",
|
||||
"DQD_DB_PREFIX": "sx_",
|
||||
"DQD_REDIS_HOST": "docker-redis",
|
||||
"DQD_REDIS_PORT": "6380",
|
||||
"DQD_REDIS_DB": "5",
|
||||
}
|
||||
with patch.dict(os.environ, env, clear=False):
|
||||
config_module._config = None
|
||||
cfg = config_module.load_config(str(settings))
|
||||
|
||||
self.assertEqual(cfg.database.host, "docker-db")
|
||||
self.assertEqual(cfg.database.port, 3307)
|
||||
self.assertEqual(cfg.database.database, "docker_name")
|
||||
self.assertEqual(cfg.database.username, "docker_user")
|
||||
self.assertEqual(cfg.database.password, "docker_password")
|
||||
self.assertEqual(cfg.database.prefix, "sx_")
|
||||
self.assertEqual(cfg.redis.host, "docker-redis")
|
||||
self.assertEqual(cfg.redis.port, 6380)
|
||||
self.assertEqual(cfg.redis.db, 5)
|
||||
|
||||
def test_render_crontab_skips_inactive_tasks_and_quotes_arguments(self):
|
||||
from scripts.docker_task_runner import CrawlerTask, render_crontab
|
||||
|
||||
tasks = [
|
||||
CrawlerTask("nba", "NBA赛程", "nba_match", "*/1 * * * *", True),
|
||||
CrawlerTask("old", "旧任务", "standings", "0 2 * * *", False),
|
||||
]
|
||||
|
||||
rendered = render_crontab(tasks, python_bin="/usr/local/bin/python", app_dir="/app")
|
||||
|
||||
self.assertIn("SHELL=/bin/bash", rendered)
|
||||
self.assertIn("BASH_ENV=/etc/sport-era-crawler/env.sh", rendered)
|
||||
self.assertIn("*/1 * * * * root cd /app && /usr/local/bin/python scripts/docker_task_runner.py run nba_match --task-key nba --task-name", rendered)
|
||||
self.assertIn("--cron-expression '*/1 * * * *'", rendered)
|
||||
self.assertNotIn("standings", rendered)
|
||||
|
||||
def test_load_tasks_accepts_top_level_list_yaml(self):
|
||||
from scripts.docker_task_runner import load_tasks
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
tasks_file = Path(tmp) / "crawler_tasks.yaml"
|
||||
tasks_file.write_text(
|
||||
"\n".join(
|
||||
[
|
||||
"- key: live",
|
||||
" name: 实时详情",
|
||||
" action: live_detail",
|
||||
" cron: '*/1 * * * *'",
|
||||
]
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
tasks = load_tasks(tasks_file)
|
||||
|
||||
self.assertEqual(len(tasks), 1)
|
||||
self.assertEqual(tasks[0].task_key, "live")
|
||||
self.assertEqual(tasks[0].action, "live_detail")
|
||||
|
||||
def test_run_task_marks_success_and_flushes_output(self):
|
||||
from scripts.docker_task_runner import CommandResult, CrawlerTask, run_task
|
||||
|
||||
store = FakeStore()
|
||||
task = CrawlerTask("nba", "NBA赛程", "nba_match", "*/1 * * * *", True, timeout=3)
|
||||
|
||||
def fake_runner(action, timeout, output_callback):
|
||||
output_callback("hello\n")
|
||||
return CommandResult(exit_code=0, output="hello\nok\n", elapsed=0.2, timed_out=False)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
result = run_task(task, store=store, lock_dir=Path(tmp), command_runner=fake_runner)
|
||||
|
||||
self.assertEqual(result.status, 1)
|
||||
self.assertEqual(store.created[0].action, "nba_match")
|
||||
self.assertEqual(store.progress[0][1], "hello\n")
|
||||
self.assertEqual(store.finished[0][1], 1)
|
||||
self.assertIn("ok", store.finished[0][2])
|
||||
|
||||
def test_run_task_marks_failure(self):
|
||||
from scripts.docker_task_runner import CommandResult, CrawlerTask, run_task
|
||||
|
||||
store = FakeStore()
|
||||
task = CrawlerTask("news", "资讯", "league_news", "*/30 * * * *", True, timeout=3)
|
||||
|
||||
def fake_runner(action, timeout, output_callback):
|
||||
return CommandResult(exit_code=7, output="boom\n", elapsed=0.1, timed_out=False)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
result = run_task(task, store=store, lock_dir=Path(tmp), command_runner=fake_runner)
|
||||
|
||||
self.assertEqual(result.status, 2)
|
||||
self.assertEqual(store.finished[0][1], 2)
|
||||
self.assertIn("退出码 7", store.finished[0][4])
|
||||
|
||||
def test_run_task_marks_timeout_as_failure(self):
|
||||
from scripts.docker_task_runner import CommandResult, CrawlerTask, run_task
|
||||
|
||||
store = FakeStore()
|
||||
task = CrawlerTask("live", "实时", "live_detail", "*/1 * * * *", True, timeout=1)
|
||||
|
||||
def fake_runner(action, timeout, output_callback):
|
||||
return CommandResult(exit_code=-1, output="slow\n", elapsed=1.0, timed_out=True)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
result = run_task(task, store=store, lock_dir=Path(tmp), command_runner=fake_runner)
|
||||
|
||||
self.assertEqual(result.status, 2)
|
||||
self.assertIn("执行超时", store.finished[0][4])
|
||||
|
||||
def test_run_task_skips_duplicate_action(self):
|
||||
from scripts.docker_task_runner import CrawlerTask, TaskLock, run_task
|
||||
|
||||
store = FakeStore()
|
||||
task = CrawlerTask("live", "实时", "live_detail", "*/1 * * * *", True, timeout=1)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
lock_dir = Path(tmp)
|
||||
lock = TaskLock("live_detail", lock_dir)
|
||||
self.assertTrue(lock.acquire())
|
||||
try:
|
||||
result = run_task(
|
||||
task,
|
||||
store=store,
|
||||
lock_dir=lock_dir,
|
||||
command_runner=lambda action, timeout, output_callback: self.fail("duplicate should not run"),
|
||||
)
|
||||
finally:
|
||||
lock.release()
|
||||
|
||||
self.assertEqual(result.status, 3)
|
||||
self.assertEqual(store.finished[0][1], 3)
|
||||
self.assertIn("已有任务执行中", store.finished[0][2])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,46 @@
|
||||
import importlib.util
|
||||
import pathlib
|
||||
import sys
|
||||
import tempfile
|
||||
import types
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = pathlib.Path(__file__).resolve().parents[1]
|
||||
sys.modules.setdefault("src.core.alert_manager", types.ModuleType("src.core.alert_manager"))
|
||||
sys.modules["src.core.alert_manager"].CrontabAlertManager = object
|
||||
sys.modules.setdefault("src.core.config", types.ModuleType("src.core.config"))
|
||||
sys.modules["src.core.config"].load_config = lambda: None
|
||||
sys.modules.setdefault("src.core.error_collector", types.ModuleType("src.core.error_collector"))
|
||||
sys.modules["src.core.error_collector"].ErrorCollector = object
|
||||
sys.modules.setdefault("src.core.logger", types.ModuleType("src.core.logger"))
|
||||
sys.modules["src.core.logger"].setup_logger = lambda: None
|
||||
SPEC = importlib.util.spec_from_file_location("crawler_main", ROOT / "main.py")
|
||||
crawler_main = importlib.util.module_from_spec(SPEC)
|
||||
SPEC.loader.exec_module(crawler_main)
|
||||
|
||||
|
||||
class RuntimeGuardTest(unittest.TestCase):
|
||||
def test_action_lock_blocks_duplicate_running_action(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
first = crawler_main.CrawlerActionLock("nba_match", pathlib.Path(tmp))
|
||||
second = crawler_main.CrawlerActionLock("nba_match", pathlib.Path(tmp))
|
||||
|
||||
self.assertTrue(first.acquire())
|
||||
try:
|
||||
self.assertFalse(second.acquire())
|
||||
finally:
|
||||
first.release()
|
||||
|
||||
self.assertTrue(second.acquire())
|
||||
second.release()
|
||||
|
||||
def test_action_timeout_uses_specific_value_or_default(self):
|
||||
self.assertEqual(crawler_main.get_action_timeout("nba_match"), 180)
|
||||
self.assertEqual(crawler_main.get_action_timeout("live_detail"), 180)
|
||||
self.assertEqual(crawler_main.get_action_timeout("article_content_fetch"), 900)
|
||||
self.assertEqual(crawler_main.get_action_timeout("unknown_task"), 600)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user