"""API v3 容灾恢复任务。 处理服务重启后的任务恢复: - 扫描处于中间状态的 ApiGenerationTask - 重新入队未完成的 Celery 任务 - 处理租约过期的任务 Worker 启动时会自动触发恢复扫描。 """ from __future__ import annotations import logging from datetime import datetime, timedelta, timezone from app.config import settings from app.enums.celery_queue import CeleryQueue from app.models.api.api_generation_task import ApiGenerationTask from app.models.base import async_session from app.tasks.api_generation_tasks import ( QUEUE_CREATE, QUEUE_DOWNLOAD, QUEUE_POLL, api_create_generation_task, api_download_generation_result_task, api_poll_generation_task, ) from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app logger = logging.getLogger("videogen") async def recover_api_generation_tasks_once(): """扫描并恢复未完成的 API v3 生成任务。 恢复场景: 1. status=pending 且未入队 -> 重新入队创建任务 2. status=generating 且 provider_task_id 为空 -> 重新入队创建任务 3. status=generating 且 provider_task_id 存在 -> 重新入队轮询任务 4. pipeline_stage=result_ready -> 重新入队下载任务 5. 租约过期但任务未完成 -> 重新入队对应阶段任务 """ now = datetime.now(timezone.utc) recovered = 0 async with async_session() as db: # 1. 恢复 pending/generating 任务(未开始或中断) result = await db.execute( __import__("sqlalchemy").select(ApiGenerationTask).where( ApiGenerationTask.status.in_(["pending", "generating"]), ApiGenerationTask.deleted_at.is_(None), ApiGenerationTask.created_at > now - timedelta(hours=48), ) ) tasks = list(result.scalars().all()) for task in tasks: try: if task.status == "pending" or not task.provider_task_id: # 重新入队创建任务 api_create_generation_task.apply_async( args=[task.id], queue=QUEUE_CREATE, ) logger.info("API recovery: re-enqueued create task %s", task.id) recovered += 1 elif task.status == "generating" and task.provider_task_id: # 检查是否需要轮询 next_poll_at = task.next_poll_at if next_poll_at is None or next_poll_at <= now: # 重新入队轮询任务 api_poll_generation_task.apply_async( args=[task.id], queue=QUEUE_POLL, ) logger.info("API recovery: re-enqueued poll task %s (provider_task_id=%s)", task.id, task.provider_task_id) recovered += 1 # 检查下载阶段 if task.pipeline_stage == "result_ready" and not task.video_url and not task.image_url: api_download_generation_result_task.apply_async( args=[task.id], queue=QUEUE_DOWNLOAD, ) logger.info("API recovery: re-enqueued download task %s", task.id) recovered += 1 except Exception as exc: logger.warning("API recovery: failed to recover task %s: %s", task.id, exc) # 2. 恢复排队任务(服务重启后,排队任务需要重新检查并发) from app.services.api_v3.quota_service import can_start_video_task, get_queued_video_tasks from app.models.api.api_key import ApiKey # 获取所有有排队任务的 API Key queued_result = await db.execute( __import__("sqlalchemy").select(ApiGenerationTask.api_key_id).where( ApiGenerationTask.status == "queued", ApiGenerationTask.deleted_at.is_(None), ApiGenerationTask.created_at > now - timedelta(hours=48), ).distinct() ) api_key_ids = [row[0] for row in queued_result.all()] for api_key_id in api_key_ids: try: key = await db.get(ApiKey, api_key_id) if not key: continue # 检查是否可以启动排队任务 if await can_start_video_task(key, db): queued_tasks = await get_queued_video_tasks(key, db, limit=1) if queued_tasks: next_task = queued_tasks[0] next_task.status = "pending" next_task.pipeline_stage = "queued" await db.commit() api_create_generation_task.apply_async( args=[next_task.id], queue=QUEUE_CREATE, ) logger.info("API recovery: started queued task %s for key %s", next_task.id, api_key_id) recovered += 1 except Exception as exc: logger.warning("API recovery: failed to recover queued task for key %s: %s", api_key_id, exc) if recovered: logger.info("API recovery: recovered %d tasks", recovered) return recovered @celery_app.task( name="api_generation.recover_tasks_once", bind=True, max_retries=0, soft_time_limit=300, time_limit=600, ) def api_generation_recover_tasks_once(self): """API v3 任务恢复扫描(Celery Beat 定时触发)。""" return run_async(recover_api_generation_tasks_once())