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, Report, Task, utc_now 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) STAGE_LABELS = { "queued": "等待启动", "crawl_hotspots": "获取热点中", "search_items": "搜索内容中", "crawl_comments": "抓取评论中", "ai_analysis": "AI 分析中", "generate_reports": "生成报告中", "success": "已完成", "failed": "失败", } 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", current_stage="queued", last_progress_at=utc_now(), ) 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]: tasks = list(session.scalars(select(Task).order_by(Task.created_at.desc()))) for task in tasks: hydrate_task_progress(session, task) return tasks def get_task(session: Session, task_id: str) -> Task | None: task = session.get(Task, task_id) if task is not None: hydrate_task_progress(session, task) return task def hydrate_task_progress(session: Session, task: Task) -> Task: task.comments_count = session.scalar(select(func.count(Comment.id)).where(Comment.task_id == task.id)) or 0 task.reports_count = session.scalar(select(func.count(Report.id)).where(Report.task_id == task.id)) or 0 task.is_demo = task.id.startswith("demo-") task.last_progress_at = task.last_progress_at or task.created_at stage_code = task.current_stage if not stage_code and task.status == "success": stage_code = "success" elif not stage_code and task.status == "failed": stage_code = "failed" elif not stage_code and task.status == "running": stage_code = "queued" task.current_stage_label = STAGE_LABELS.get(stage_code or "", stage_code or "等待启动") return task def update_task_stage(task: Task, stage: str) -> None: task.current_stage = stage task.last_progress_at = utc_now() 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: update_task_stage(task, "crawl_hotspots") session.commit() 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: update_task_stage(task, "search_items") session.commit() 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: update_task_stage(task, "crawl_comments") session.commit() 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))) update_task_stage(task, "ai_analysis") session.commit() 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 update_task_stage(task, "generate_reports") 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: update_task_stage(task, "generate_reports") 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" update_task_stage(task, "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" update_task_stage(task, "failed") task.error_stage = stage task.error_type = error_type task.error_message = message