150 lines
6.8 KiB
Python
150 lines
6.8 KiB
Python
import json
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.models import Comment, ContentItem, Hotspot, Report, Task
|
|
from app.schemas import CreateTaskRequest
|
|
from app.services.report_service import DEFAULT_SUMMARY
|
|
from app.services.task_service import create_task, run_task
|
|
from tests.helpers import make_test_client
|
|
|
|
|
|
class FakeXhsPlatform:
|
|
def fetch_hotspots(self, *, limit):
|
|
from app.platforms.base import HotspotData
|
|
|
|
return [HotspotData(source_hot_id="h1", title="热点一", rank=1, heat_value="100", raw_data={"id": "h1"})]
|
|
|
|
def search_items_by_hotspot(self, keyword, *, limit):
|
|
from app.platforms.base import ContentItemData
|
|
|
|
return [
|
|
ContentItemData(source_item_id="n1", item_type="note", title=f"{keyword} 笔记", raw_data={"id": "n1"})
|
|
]
|
|
|
|
def fetch_comments(self, source_item_id, *, limit):
|
|
from app.platforms.base import CommentData
|
|
|
|
return [CommentData(source_comment_id="c1", content="评论一", like_count=5, raw_data={"id": "c1"})]
|
|
|
|
|
|
ai_prompts: list[str] = []
|
|
|
|
|
|
def successful_ai_response(prompt: str) -> str:
|
|
ai_prompts.append(prompt)
|
|
comments = json.loads(prompt[prompt.index("[") :])
|
|
comment_id = comments[0]["comment_id"]
|
|
return json.dumps(
|
|
[{"comment_id": comment_id, "sentiment": "positive", "labels": ["认可"], "reason": "喜欢"}],
|
|
ensure_ascii=False,
|
|
)
|
|
|
|
|
|
def test_xiaohongshu_task_flow_persists_hotspots_items_and_comments(monkeypatch):
|
|
ai_prompts.clear()
|
|
with make_test_client() as (_client, engine):
|
|
monkeypatch.setattr("app.services.task_service.build_platform", lambda _platform: FakeXhsPlatform())
|
|
summary_provider = lambda metrics, typical, *, word_limit=200: (
|
|
ai_prompts.append(f"只返回一段中文总结 word_limit={word_limit} sample={metrics['sample_count']}")
|
|
or "这是一段真实 AI 报告摘要。"
|
|
)
|
|
monkeypatch.setattr("app.services.task_service.build_ai_dependencies", lambda: (successful_ai_response, summary_provider))
|
|
|
|
with Session(engine) as session:
|
|
task = create_task(
|
|
session,
|
|
CreateTaskRequest(platform="xiaohongshu", hotspot_limit=1, item_limit_per_hotspot=1, comment_limit_per_item=10),
|
|
submit_background=False,
|
|
)
|
|
task_id = task.id
|
|
|
|
with Session(engine) as session:
|
|
run_task(task_id, session_factory=lambda: session)
|
|
|
|
persisted = session.get(Task, task_id)
|
|
assert persisted.status == "success"
|
|
assert persisted.processed_items_count == 1
|
|
assert persisted.successful_items_count == 1
|
|
assert persisted.analysis_status == "normal"
|
|
assert session.scalar(select(Hotspot).where(Hotspot.task_id == task_id)).title == "热点一"
|
|
assert session.scalar(select(ContentItem).where(ContentItem.task_id == task_id)).source_item_id == "n1"
|
|
comment = session.scalar(select(Comment).where(Comment.task_id == task_id))
|
|
assert comment.content == "评论一"
|
|
assert comment.ai_analysis_status == "success"
|
|
assert comment.sentiment == "positive"
|
|
assert comment.labels == '["认可"]'
|
|
assert comment.reason == "喜欢"
|
|
item_report = session.scalar(select(Report).where(Report.task_id == task_id, Report.report_type == "item"))
|
|
hotspot_report = session.scalar(select(Report).where(Report.task_id == task_id, Report.report_type == "hotspot"))
|
|
assert item_report is not None
|
|
assert hotspot_report is not None
|
|
assert item_report.summary == "这是一段真实 AI 报告摘要。"
|
|
assert hotspot_report.summary == "这是一段真实 AI 报告摘要。"
|
|
assert DEFAULT_SUMMARY not in item_report.markdown_content
|
|
assert DEFAULT_SUMMARY not in hotspot_report.markdown_content
|
|
assert any("只返回一段中文总结" in prompt for prompt in ai_prompts)
|
|
|
|
|
|
def test_xiaohongshu_task_flow_marks_comments_failed_when_ai_parse_fails(monkeypatch):
|
|
with make_test_client() as (_client, engine):
|
|
monkeypatch.setattr("app.services.task_service.build_platform", lambda _platform: FakeXhsPlatform())
|
|
monkeypatch.setattr("app.services.task_service.build_ai_dependencies", lambda: (lambda _prompt: "not-json", None))
|
|
monkeypatch.setattr("app.services.ai_service.time.sleep", lambda _seconds: None)
|
|
|
|
with Session(engine) as session:
|
|
task = create_task(
|
|
session,
|
|
CreateTaskRequest(platform="xiaohongshu", hotspot_limit=1, item_limit_per_hotspot=1, comment_limit_per_item=10),
|
|
submit_background=False,
|
|
)
|
|
task_id = task.id
|
|
|
|
with Session(engine) as session:
|
|
run_task(task_id, session_factory=lambda: session)
|
|
|
|
persisted = session.get(Task, task_id)
|
|
comment = session.scalar(select(Comment).where(Comment.task_id == task_id))
|
|
|
|
assert persisted.status == "success"
|
|
assert persisted.analysis_status == "insufficient"
|
|
assert persisted.analysis_success_rate == 0.0
|
|
assert comment.ai_analysis_status == "failed"
|
|
assert comment.sentiment == "unknown"
|
|
assert comment.labels == "[]"
|
|
assert comment.reason == "ai_parse_failed"
|
|
|
|
|
|
def test_xiaohongshu_task_flow_keeps_item_success_when_ai_request_raises(monkeypatch):
|
|
with make_test_client() as (_client, engine):
|
|
monkeypatch.setattr("app.services.task_service.build_platform", lambda _platform: FakeXhsPlatform())
|
|
monkeypatch.setattr("app.services.ai_service.time.sleep", lambda _seconds: None)
|
|
|
|
def failing_requester(_prompt):
|
|
raise RuntimeError("ai unauthorized")
|
|
|
|
monkeypatch.setattr("app.services.task_service.build_ai_dependencies", lambda: (failing_requester, None))
|
|
|
|
with Session(engine) as session:
|
|
task = create_task(
|
|
session,
|
|
CreateTaskRequest(platform="xiaohongshu", hotspot_limit=1, item_limit_per_hotspot=1, comment_limit_per_item=10),
|
|
submit_background=False,
|
|
)
|
|
task_id = task.id
|
|
|
|
with Session(engine) as session:
|
|
run_task(task_id, session_factory=lambda: session)
|
|
|
|
persisted = session.get(Task, task_id)
|
|
comment = session.scalar(select(Comment).where(Comment.task_id == task_id))
|
|
|
|
assert persisted.status == "success"
|
|
assert persisted.successful_items_count == 1
|
|
assert persisted.failed_items_count == 0
|
|
assert persisted.analysis_status == "insufficient"
|
|
assert persisted.analysis_success_rate == 0.0
|
|
assert comment.ai_analysis_status == "failed"
|
|
assert comment.reason == "ai_parse_failed"
|