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