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, 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_requester(): settings = get_settings() if settings.ai_base_url and settings.ai_api_key and settings.ai_model: client = OpenAICompatibleAIClient( base_url=settings.ai_base_url, api_key=settings.ai_api_key, model=settings.ai_model, timeout_seconds=settings.ai_timeout_seconds, ) return client.request def fallback_requester(_prompt: str) -> str: return "[]" return fallback_requester 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) 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=build_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) 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) 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