from __future__ import annotations import logging from datetime import datetime, timedelta, timezone from typing import Any, Dict from sqlalchemy import or_, select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.models.chat_generation_task import ChatGenerationTask from app.services.celery_download_recovery_service import ( ensure_aware_utc, get_download_active_payloads, get_due_download_record_ids, postpone_download_active_check, remove_download_active, ) from app.services.generation_log_service import log_task_event from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once logger = logging.getLogger("video_gen") def _now() -> datetime: return datetime.now(timezone.utc) def _is_expired(value: datetime | None, now: datetime | None = None) -> bool: checked = ensure_aware_utc(value) if checked is None: return True return checked <= (now or _now()) def _queue_timeout_at(task: ChatGenerationTask, now: datetime | None = None) -> datetime: current_time = now or _now() enqueued_at = ensure_aware_utc(task.download_enqueued_at) if enqueued_at is None: return current_time return enqueued_at + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300)) def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) -> bool: current_time = now or _now() return _queue_timeout_at(task, current_time) <= current_time def _is_final_task_state(task: ChatGenerationTask) -> bool: return task.status in ("completed", "failed") or task.pipeline_stage in ( "done", "failed", "timeout", "download_failed", ) async def recover_one_download_task( db: AsyncSession, task: ChatGenerationTask, *, payload: dict[str, Any] | None = None, source: str = "startup_db", ) -> str: from app.tasks.generation_download_tasks import ( DOWNLOAD_STAGE_DOWNLOADING, DOWNLOAD_STAGE_QUEUED, DOWNLOAD_STAGE_RETRY_WAITING, enqueue_download_task, ) current_time = _now() if not task: return "skip_missing_task" if task.generation_mode != "chatapi_async": await remove_download_active(task.id) return "clean_invalid_mode" if _is_final_task_state(task): await remove_download_active(task.id) return "clean_final_state" if task.status != "generating": await remove_download_active(task.id) return "clean_not_generating" if not task.remote_result_url: return "skip_no_remote_result_url" stage = task.pipeline_stage redis_payload = payload or {} if stage == "result_ready": await log_task_event( task, event_type="DOWNLOAD_RECOVERY_ENQUEUE", message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) await enqueue_download_task( db, task, recover=True, reason=f"{source}_result_ready", ) return "recover_result_ready" if stage == DOWNLOAD_STAGE_QUEUED: if _is_queue_timeout(task, current_time): await log_task_event( task, event_type="DOWNLOAD_RECOVERY_ENQUEUE", message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) await enqueue_download_task( db, task, recover=True, reason=f"{source}_download_queued_timeout", ) return "recover_queued_timeout" await postpone_download_active_check( record_id=task.id, payload=payload, check_at=_queue_timeout_at(task, current_time), ) return "skip_queued_not_timeout" if stage == DOWNLOAD_STAGE_DOWNLOADING: if _is_expired(task.download_lease_until, current_time): await log_task_event( task, event_type="DOWNLOAD_RECOVERY_ENQUEUE", message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) await enqueue_download_task( db, task, recover=True, reason=f"{source}_downloading_lease_expired", ) return "recover_downloading_expired" await postpone_download_active_check( record_id=task.id, payload=payload, check_at=task.download_lease_until, ) return "skip_downloading_alive" if stage == DOWNLOAD_STAGE_RETRY_WAITING: if _is_expired(task.download_next_retry_at, current_time): await log_task_event( task, event_type="DOWNLOAD_RECOVERY_ENQUEUE", message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) await enqueue_download_task( db, task, recover=True, reason=f"{source}_retry_waiting_due", ) return "recover_retry_due" await postpone_download_active_check( record_id=task.id, payload=payload, check_at=task.download_next_retry_at, ) return "skip_retry_waiting_not_due" return f"skip_stage_{stage}" async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]: """启动时下载容灾扫描。 先按 Redis active_index 找到到期下载任务;Redis 不可用或索引丢失时, 再通过 DB fallback 扫描 result_ready/download_* 状态,避免任务永久卡住。 """ checked_ids: set[str] = set() results: dict[str, int] = {} due_ids = await get_due_download_record_ids( limit=settings.DOWNLOAD_RECOVERY_BATCH_SIZE, ) payloads = await get_download_active_payloads(due_ids) for task_id in due_ids: result = await db.execute( select(ChatGenerationTask) .where( ChatGenerationTask.id == task_id, ChatGenerationTask.deleted_at.is_(None), ) .with_for_update() .limit(1) ) task = result.scalar_one_or_none() if task is None: await remove_download_active(task_id) action = "clean_missing_task" else: checked_ids.add(task.id) action = await recover_one_download_task( db, task, payload=payloads.get(task_id), source="startup_redis", ) results[action] = results.get(action, 0) + 1 # DB fallback:不依赖 Redis active 注册表。 fallback_result = await db.execute( select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode == "chatapi_async", ChatGenerationTask.status == "generating", ChatGenerationTask.remote_result_url.is_not(None), ChatGenerationTask.pipeline_stage.in_( ["result_ready", "download_queued", "downloading", "retry_waiting"] ), ) .order_by(ChatGenerationTask.updated_at.asc()) .limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)) .with_for_update(skip_locked=True) ) fallback_tasks = fallback_result.scalars().all() for task in fallback_tasks: if task.id in checked_ids: continue action = await recover_one_download_task( db, task, payload=None, source="startup_db", ) results[action] = results.get(action, 0) + 1 checked_ids.add(task.id) return {"checked": len(checked_ids), "results": results} async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: """启动时生成链路容灾扫描。 只在 Celery worker 启动时跑一次,不引入 beat,不新增第四条启动命令。 用于把 queued/creating/waiting_remote/polling/result_ready 等中间态重新投递到现有三个队列。 """ from app.tasks.generation_create_tasks import chatapi_create_generation_task from app.tasks.generation_download_tasks import enqueue_download_task from app.tasks.generation_poll_tasks import poll_generation_task current_time = _now() results: dict[str, int] = {} query_result = await db.execute( select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode == "chatapi_async", ChatGenerationTask.status == "generating", ChatGenerationTask.pipeline_stage.in_( [ "queued", "preparing", "creating_provider_task", "waiting_remote", "polling", "result_ready", ] ), ) .order_by(ChatGenerationTask.updated_at.asc()) .limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)) .with_for_update(skip_locked=True) ) tasks = query_result.scalars().all() for task in tasks: if task.deadline_at and _is_expired(task.deadline_at, current_time): await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message="任务超时", pipeline_stage="timeout", ) await db.commit() await log_task_event( task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout", ) action = "mark_timeout" elif task.pipeline_stage in ("queued", "preparing", "creating_provider_task"): await log_task_event( task, event_type="GENERATION_RECOVERY_ENQUEUE", message="启动时发现创建阶段任务未完成,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage}, ) chatapi_create_generation_task.apply_async( args=[task.id], queue="gen_chatapi_create", countdown=0, ) action = "recover_create" elif task.pipeline_stage in ("waiting_remote", "polling"): if task.remote_result_url: await enqueue_download_task( db, task, recover=True, reason="startup_waiting_remote_has_result", ) action = "recover_waiting_has_result" elif task.provider_task_id or task.seedance_task_id: await log_task_event( task, event_type="GENERATION_RECOVERY_ENQUEUE", message="启动时发现远程等待/轮询阶段任务未完成,恢复投递轮询队列", detail={"pipeline_stage": task.pipeline_stage}, ) task.pipeline_stage = "waiting_remote" await db.commit() poll_generation_task.apply_async( args=[task.id], queue="gen_provider_poll", countdown=0, ) action = "recover_poll" else: await log_task_event( task, event_type="GENERATION_RECOVERY_ENQUEUE", message="启动时发现任务缺少供应商任务ID,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage}, ) task.pipeline_stage = "queued" await db.commit() chatapi_create_generation_task.apply_async( args=[task.id], queue="gen_chatapi_create", countdown=0, ) action = "recover_create_missing_provider_id" elif task.pipeline_stage == "result_ready": if task.remote_result_url: await enqueue_download_task( db, task, recover=True, reason="startup_generation_result_ready", ) action = "recover_result_ready" else: action = "skip_result_ready_no_url" else: action = f"skip_stage_{task.pipeline_stage}" results[action] = results.get(action, 0) + 1 # 下载阶段单独跑 DB fallback。 download_result = await recover_download_tasks_once(db) return { "checked": len(tasks), "results": results, "download_recovery": download_result, }