迁移目录

This commit is contained in:
hajimi
2026-06-12 22:27:18 +08:00
parent 7381097582
commit cd44cd6e47
92 changed files with 1839 additions and 0 deletions
+1
View File
@@ -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()
+111
View File
@@ -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()