Files
hot_comment_radar/app/services/task_service.py
T

274 lines
12 KiB
Python

import inspect
import json
from concurrent.futures import ThreadPoolExecutor
from collections.abc import Callable
from sqlalchemy import func, select
from sqlalchemy.orm import Session, sessionmaker
from app.config import get_settings
from app.db import SessionLocal
from app.models import Comment, ContentItem, Hotspot, Task
from app.platforms.base import PlatformAPIError, TikHubClient
from app.platforms.douyin import DouyinPlatform
from app.platforms.xiaohongshu import XiaohongshuPlatform
from app.schemas import CreateTaskRequest
from app.services.ai_service import OpenAICompatibleAIClient, analyze_comments_with_retry, build_report_summary_prompt, calculate_analysis_status
from app.services.report_service import generate_hotspot_report, generate_item_report
RUNNING_TASK_MESSAGE = "当前有正在运行的任务,请稍后再试"
RESTART_ERROR_MESSAGE = "系统重启,任务被中断"
task_executor = ThreadPoolExecutor(max_workers=1)
def has_running_task(session: Session) -> bool:
return session.scalar(select(Task.id).where(Task.status == "running").limit(1)) is not None
def recover_running_tasks(session: Session) -> int:
tasks = list(session.scalars(select(Task).where(Task.status == "running")))
for task in tasks:
task.status = "failed"
task.error_stage = "system"
task.error_type = "unexpected_restart"
task.error_message = RESTART_ERROR_MESSAGE
session.commit()
return len(tasks)
def build_platform(platform: str):
settings = get_settings()
client = TikHubClient(
base_url=settings.tikhub_base_url,
api_key=settings.tikhub_api_key,
timeout_seconds=settings.http_timeout_seconds,
max_retries=settings.http_max_retries,
)
if platform == "xiaohongshu":
return XiaohongshuPlatform(client)
if platform == "douyin":
return DouyinPlatform(client)
raise ValueError(f"Unsupported platform: {platform}")
def build_ai_client() -> OpenAICompatibleAIClient | None:
settings = get_settings()
if settings.ai_base_url and settings.ai_api_key and settings.ai_model:
return OpenAICompatibleAIClient(
base_url=settings.ai_base_url,
api_key=settings.ai_api_key,
model=settings.ai_model,
timeout_seconds=settings.ai_timeout_seconds,
)
return None
def build_ai_requester(client: OpenAICompatibleAIClient | None = None):
client = client if client is not None else build_ai_client()
if client is not None:
return lambda prompt: client.request(
prompt,
system_prompt="只返回 JSON Array,不输出 Markdown 或解释性自然语言。",
)
def fallback_requester(_prompt: str) -> str:
return "[]"
return fallback_requester
def build_report_summary_provider(client: OpenAICompatibleAIClient | None):
if client is None:
return None
def summarize_report(metrics: dict, typical: dict, *, word_limit: int = 200) -> str:
return client.request(
build_report_summary_prompt(metrics, typical, word_limit=word_limit),
system_prompt="只返回一段中文总结,不使用 Markdown。",
)
return summarize_report
def build_ai_dependencies():
client = build_ai_client()
requester = build_ai_requester(client)
return requester, build_report_summary_provider(client)
def create_task(session: Session, request: CreateTaskRequest, *, submit_background: bool = True) -> Task:
task = Task(
platform=request.platform,
status="running",
hotspot_limit=request.hotspot_limit,
item_limit_per_hotspot=request.item_limit_per_hotspot,
comment_limit_per_item=request.comment_limit_per_item,
total_items_count=0,
processed_items_count=0,
successful_items_count=0,
failed_items_count=0,
analysis_success_rate=0.0,
analysis_status="normal",
)
session.add(task)
session.commit()
session.refresh(task)
if submit_background:
task_executor.submit(run_task, task.id)
return task
def list_tasks(session: Session) -> list[Task]:
return list(session.scalars(select(Task).order_by(Task.created_at.desc())))
def get_task(session: Session, task_id: str) -> Task | None:
return session.get(Task, task_id)
def run_task(task_id: str, *, session_factory: Callable[[], Session] | sessionmaker = SessionLocal) -> None:
session = session_factory()
close_session = hasattr(session, "close")
try:
task = session.get(Task, task_id)
if task is None:
return
try:
platform = build_platform(task.platform)
ai_requester, report_summary_provider = build_ai_dependencies()
try:
hotspots = platform.fetch_hotspots(limit=task.hotspot_limit)
except PlatformAPIError as exc:
_mark_task_failed(task, "crawl_hotspots", exc.error_type, str(exc))
session.commit()
return
except Exception as exc:
_mark_task_failed(task, "crawl_hotspots", "api_error", str(exc))
session.commit()
return
for hotspot_data in hotspots:
hotspot = Hotspot(
task_id=task.id,
platform=task.platform,
source_hot_id=hotspot_data.source_hot_id,
rank=hotspot_data.rank,
title=hotspot_data.title,
heat_value=hotspot_data.heat_value,
raw_data=json.dumps(hotspot_data.raw_data or {}, ensure_ascii=False),
)
session.add(hotspot)
session.flush()
try:
items = platform.search_items_by_hotspot(hotspot.title, limit=task.item_limit_per_hotspot)
except Exception as exc:
task.failed_items_count += task.item_limit_per_hotspot
task.error_stage = "crawl_items"
task.error_type = getattr(exc, "error_type", "api_error")
task.error_message = str(exc)
session.commit()
continue
task.total_items_count += len(items)
seen_item_ids: set[str] = set()
for item_data in items:
if item_data.source_item_id in seen_item_ids:
continue
seen_item_ids.add(item_data.source_item_id)
item = ContentItem(
task_id=task.id,
hotspot_id=hotspot.id,
platform=task.platform,
source_item_id=item_data.source_item_id,
item_type=item_data.item_type,
title=item_data.title,
summary=item_data.summary,
url=item_data.url,
status="pending",
raw_data=json.dumps(item_data.raw_data or {}, ensure_ascii=False),
)
session.add(item)
session.flush()
try:
comments = platform.fetch_comments(item.source_item_id, limit=task.comment_limit_per_item)
for comment_data in comments:
session.add(
Comment(
task_id=task.id,
hotspot_id=hotspot.id,
content_item_id=item.id,
platform=task.platform,
source_comment_id=comment_data.source_comment_id,
content=comment_data.content,
author=comment_data.author,
like_count=comment_data.like_count,
comment_time=comment_data.comment_time,
raw_data=json.dumps(comment_data.raw_data or {}, ensure_ascii=False),
)
)
session.flush()
item_comments = list(session.scalars(select(Comment).where(Comment.content_item_id == item.id)))
ai_results = analyze_comments_with_retry(
[{"comment_id": comment.id, "content": comment.content} for comment in item_comments],
requester=ai_requester,
max_retries=get_settings().ai_max_retries,
)
results_by_comment_id = {result.comment_id: result for result in ai_results}
for comment in item_comments:
result = results_by_comment_id.get(comment.id)
if result is None:
comment.sentiment = "unknown"
comment.labels = "[]"
comment.reason = "ai_missing_result"
comment.ai_analysis_status = "failed"
continue
comment.sentiment = result.sentiment
comment.labels = json.dumps(result.labels, ensure_ascii=False)
comment.reason = result.reason
comment.ai_analysis_status = result.ai_analysis_status
item.status = "success"
task.successful_items_count += 1
generate_item_report(session, item.id, summary_provider=report_summary_provider)
except Exception as exc:
item.status = "failed"
item.error_stage = "crawl_comments"
item.error_type = getattr(exc, "error_type", "api_error")
item.error_message = str(exc)
task.failed_items_count += 1
task.error_stage = item.error_stage
task.error_type = item.error_type
task.error_message = item.error_message
finally:
task.processed_items_count += 1
session.commit()
if task.successful_items_count > 0:
total_comments = session.scalar(select(func.count(Comment.id)).where(Comment.task_id == task.id)) or 0
success_comments = sum(1 for comment in task.comments if comment.ai_analysis_status == "success")
task.analysis_success_rate, task.analysis_status = calculate_analysis_status(
success_count=success_comments,
total_count=total_comments,
)
for hotspot in task.hotspots:
generate_hotspot_report(session, hotspot.id, summary_provider=report_summary_provider)
task.status = "success"
else:
_mark_task_failed(task, task.error_stage or "crawl_items", task.error_type or "no_successful_items", task.error_message or "没有任何内容条目成功")
session.commit()
except Exception as exc:
_mark_task_failed(task, "system", "unexpected_error", str(exc))
session.commit()
finally:
if close_session:
session.close()
def _mark_task_failed(task: Task, stage: str, error_type: str, message: str) -> None:
task.status = "failed"
task.error_stage = stage
task.error_type = error_type
task.error_message = message