from __future__ import annotations import json import logging from datetime import datetime, timedelta, timezone from typing import Any from sqlalchemy import or_, select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.enums.celery_queue import CeleryQueue from app.enums.generation_task import ( ALLOWED_GENERATION_MODES, ChatGenerationPipelineStage, ChatGenerationTaskEventType, ChatGenerationTaskStatus, GenerationMode, GenerationOwnerType, GenerationType, ) from app.models.chat_generation_task import ChatGenerationTask from app.services.celery_runtime.runtime_service import runtime_lock_exists, runtime_lock_values 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.pipeline.db_lock_service import apply_short_lock_timeout from app.services.generation.pipeline.lifecycle_service import notify_owner_finished from app.services.generation.poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.pipeline.owner_service import ( load_generation_owner, parse_redis_owner_item_id, redis_owner_item_id, ) from app.services.redis_registry_service import ( redis_get_due_registry_ids, redis_get_registry_payloads, redis_postpone_registry_item, redis_remove_registry_item, ) logger = logging.getLogger("video_gen") POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value def _create_lock_key(task: ChatGenerationTask) -> str: return ( f"{settings.GENERATION_CREATE_LOCK_KEY_PREFIX}:" f"{GenerationOwnerType.CHAT_GENERATION_TASK.value}:{task.id}:" f"attempt:{int(task.generation_attempt_no or 1)}" ) def _poll_lock_key(task: ChatGenerationTask) -> str: return ( f"{settings.GENERATION_POLL_LOCK_KEY_PREFIX}:" f"{GenerationOwnerType.CHAT_GENERATION_TASK.value}:{task.id}:" f"attempt:{int(task.generation_attempt_no or 1)}" ) def _download_lock_key(task: ChatGenerationTask) -> str: return ( f"{settings.GENERATION_DOWNLOAD_LOCK_KEY_PREFIX}:" f"{GenerationOwnerType.CHAT_GENERATION_TASK.value}:{task.id}:" f"attempt:{int(task.generation_attempt_no or 1)}" ) def _chat_registry_id(task: ChatGenerationTask) -> str: return redis_owner_item_id( GenerationOwnerType.CHAT_GENERATION_TASK.value, str(task.id), int(getattr(task, "generation_attempt_no", 1) or 1), ) async def _load_chat_task_for_update( db: AsyncSession, task_id: str ) -> ChatGenerationTask | None: owner = await load_generation_owner( db, owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value, owner_id=str(task_id), for_update=True, ) return owner if isinstance(owner, ChatGenerationTask) else None 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 (ChatGenerationTaskStatus.COMPLETED.value, ChatGenerationTaskStatus.FAILED.value) or task.pipeline_stage in ( ChatGenerationPipelineStage.DONE.value, ChatGenerationPipelineStage.FAILED.value, ChatGenerationPipelineStage.TIMEOUT.value, ChatGenerationPipelineStage.DOWNLOAD_FAILED.value, ) def _is_success(status: str | None) -> bool: return str(status or "").lower() in ("succeeded", "success", "completed", "done") def _is_failed(status: str | None) -> bool: return str(status or "").lower() in ("failed", "error", "canceled", "cancelled") def _engine_snapshot(task: ChatGenerationTask) -> dict[str, Any]: try: value = json.loads(task.engine_snapshot_json or "{}") return value if isinstance(value, dict) else {} except Exception: return {} def _poll_queue_timeout_at(now: datetime | None = None) -> datetime: current_time = now or _now() return current_time + timedelta(seconds=int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120)) async def _remove_poll_active(task_id: str) -> None: await redis_remove_registry_item( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, item_id=task_id, log_context="poll_active", ) async def _postpone_poll_active( *, task_id: str, payload: dict[str, Any] | None = None, check_at: datetime | int | float | None = None, ) -> None: await redis_postpone_registry_item( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, item_id=task_id, payload=payload, check_at=check_at or _poll_queue_timeout_at(), log_context="poll_active", ) 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 not in ALLOWED_GENERATION_MODES: await remove_download_active(_chat_registry_id(task)) return "clean_invalid_mode" if _is_final_task_state(task): await remove_download_active(_chat_registry_id(task)) return "clean_final_state" if task.status != ChatGenerationTaskStatus.GENERATING.value: task_id = str(task.id) generation_attempt_no = int(task.generation_attempt_no or 1) generation_mode = str(task.generation_mode or "") status = task.status stage = task.pipeline_stage registry_id = _chat_registry_id(task) await db.rollback() await remove_download_active(registry_id) await log_task_event( owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value, owner_id=task_id, task_id=task_id, generation_attempt_no=generation_attempt_no, generation_mode=generation_mode, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value, message=f"{source} 下载恢复跳过:任务不是 generating", detail={"status": status, "stage": stage}, ) return "clean_not_generating" if not task.remote_result_url: task_id = str(task.id) generation_attempt_no = int(task.generation_attempt_no or 1) generation_mode = str(task.generation_mode or "") status = task.status stage = task.pipeline_stage await db.rollback() await log_task_event( owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value, owner_id=task_id, task_id=task_id, generation_attempt_no=generation_attempt_no, generation_mode=generation_mode, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value, message=f"{source} 下载恢复跳过:缺少 remote_result_url", detail={"status": status, "stage": stage}, ) return "skip_no_remote_result_url" if await runtime_lock_exists(_download_lock_key(task)): return "skip_live_download_runtime_lock" stage = task.pipeline_stage redis_payload = payload or {} if stage == ChatGenerationPipelineStage.RESULT_READY.value: await enqueue_download_task( db, task, recover=True, reason=f"{source}_result_ready", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) return "recover_result_ready" if stage == DOWNLOAD_STAGE_QUEUED: if _is_queue_timeout(task, current_time): await enqueue_download_task( db, task, recover=True, reason=f"{source}_download_queued_timeout", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) return "recover_queued_timeout" await postpone_download_active_check( record_id=_chat_registry_id(task), 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 enqueue_download_task( db, task, recover=True, reason=f"{source}_downloading_lease_expired", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) return "recover_downloading_expired" await postpone_download_active_check( record_id=_chat_registry_id(task), 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 enqueue_download_task( db, task, recover=True, reason=f"{source}_retry_waiting_due", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) return "recover_retry_due" await postpone_download_active_check( record_id=_chat_registry_id(task), 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) download_refs = { registry_item_id: parse_redis_owner_item_id(registry_item_id) for registry_item_id in due_ids } for registry_item_id, ref in download_refs.items(): if ref.owner_type != GenerationOwnerType.CHAT_GENERATION_TASK.value: continue task = await _load_chat_task_for_update(db, ref.owner_id) if task is None: await db.rollback() await remove_download_active(registry_item_id) action = "clean_missing_task" elif ( ref.generation_attempt_no is not None and int(task.generation_attempt_no or 1) != int(ref.generation_attempt_no) ): await db.rollback() await remove_download_active(registry_item_id) action = "clean_stale_download_attempt" else: current_registry_id = _chat_registry_id(task) if registry_item_id != current_registry_id: await remove_download_active(registry_item_id) checked_ids.add(task.id) action = await recover_one_download_task( db, task, payload=payloads.get(registry_item_id), source="startup_redis", ) results[action] = results.get(action, 0) + 1 # DB fallback:不依赖 Redis active 注册表。 fallback_result = await db.execute( select(ChatGenerationTask.id) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.remote_result_url.is_not(None), ChatGenerationTask.pipeline_stage.in_( [ ChatGenerationPipelineStage.RESULT_READY.value, ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value, ChatGenerationPipelineStage.DOWNLOADING.value, ChatGenerationPipelineStage.RETRY_WAITING.value, ] ), ) .order_by(ChatGenerationTask.updated_at.asc()) .limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)) ) fallback_ids = [str(value) for value in fallback_result.scalars().all()] for task_id in fallback_ids: if task_id in checked_ids: continue task = await _load_chat_task_for_update(db, task_id) if task is None: await db.rollback() 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 _mark_timeout( db: AsyncSession, task: ChatGenerationTask, *, error_message: str = "任务超时", ) -> str: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=error_message, pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) await db.commit() await notify_owner_finished(db, task) await _remove_poll_active(_chat_registry_id(task)) await log_task_event( task, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, to_status="failed", to_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) return "mark_timeout" async def _mark_failed( db: AsyncSession, task: ChatGenerationTask, *, error_message: str, event_type: str = "POLL_FAILED", detail: Any = None, ) -> str: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=error_message, pipeline_stage=ChatGenerationPipelineStage.FAILED.value, ) await db.commit() await notify_owner_finished(db, task) await _remove_poll_active(_chat_registry_id(task)) await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail) return "mark_failed" async def recover_one_generation_task( db: AsyncSession, task: ChatGenerationTask, *, payload: dict[str, Any] | None = None, source: str = "startup_db", ) -> str: """恢复单个生成任务。 分流原则: 1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。 2. 已有 provider_task_id/seedance_task_id:恢复 poll。 3. 无结果 URL、无供应商任务 ID:deadline 未过才恢复 create。 4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成。 """ 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, register_poll_active current_time = _now() redis_payload = payload or {} if not task: return "skip_missing_task" if task.generation_mode not in ALLOWED_GENERATION_MODES: await _remove_poll_active(_chat_registry_id(task)) return "clean_invalid_mode" if _is_final_task_state(task): await _remove_poll_active(_chat_registry_id(task)) return "clean_final_state" if task.status != ChatGenerationTaskStatus.GENERATING.value: await _remove_poll_active(_chat_registry_id(task)) return "clean_not_generating" if bool(getattr(task, "video_upscale_enabled_snapshot", False)) and str(task.pipeline_stage or "").startswith("upscale_"): # 原视频已经进入超分流水线,后续由 video_upscale 恢复扫描处理。 # 这里禁止再次投递原结果下载,避免覆盖保留的 source.mp4 或提前生成用户资源。 await _remove_poll_active(_chat_registry_id(task)) return "delegate_video_upscale_recovery" has_remote_result = bool(str(task.remote_result_url or "").strip()) has_provider_task_id = bool(str(task.provider_task_id or "").strip() or str(task.seedance_task_id or "").strip()) is_deadline_expired = bool(task.deadline_at and _is_expired(task.deadline_at, current_time)) runtime_lock_key = ( _download_lock_key(task) if has_remote_result else (_poll_lock_key(task) if has_provider_task_id else _create_lock_key(task)) ) if await runtime_lock_exists(runtime_lock_key): return "skip_live_runtime_lock" # 最高优先级:只要远程结果 URL 已经落库,说明生成侧已经成功。 # 不管当前 pipeline_stage 是 queued/creating/waiting/result_ready/download_*,恢复时都不能重复 create 或 poll。 if has_remote_result: await _remove_poll_active(_chat_registry_id(task)) await enqueue_download_task( db, task, recover=True, reason=f"{source}_has_remote_result_url", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列", detail={ "pipeline_stage": task.pipeline_stage, "payload": redis_payload, "deadline_expired": is_deadline_expired, }, ) return "recover_download_has_remote_result" # 已经过 deadline 且没有结果 URL: # - 有供应商任务 ID:交给 poll worker 做最后一次状态确认; # - 没有供应商任务 ID:说明没有可查询的远程任务,直接按超时失败处理,不再重新 create。 if is_deadline_expired: if has_provider_task_id: task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) poll_generation_task.apply_async( args=[task.id], kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, queue=POLL_QUEUE, countdown=0, ) await register_poll_active( task, check_at=_poll_queue_timeout_at(), reason=f"{source}_deadline_final_poll", ) return "recover_deadline_final_poll" return await _mark_timeout( db, task, error_message=f"{source} 发现任务已到 deadline,且没有远程结果或供应商任务ID", ) # 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。 # 视频任务如果 next_poll_at 未到期,不提前 poll,只刷新 active 注册表等待 Beat dispatcher 到期投递。 if has_provider_task_id: if is_video_generation_task(task): ensure_video_poll_fields(task, now=current_time) if is_poll_not_due(task, now=current_time): await db.commit() await register_poll_active( task, check_at=task.next_poll_at, next_poll_at=task.next_poll_at, reason=f"{source}_video_poll_not_due", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value, message=f"{source} 发现视频任务尚未到下一次轮询时间,启动容灾不提前投递 poll", detail={ "pipeline_stage": task.pipeline_stage, "payload": redis_payload, "next_poll_at": task.next_poll_at, }, ) return "skip_video_poll_not_due" original_next_poll_at = ensure_aware_utc(task.next_poll_at) queue_hold_until = _poll_queue_timeout_at(current_time) task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value # 这里仍复用 next_poll_at 做短暂队列保护,避免启动容灾重复投递。 # 真正消费时通过 force_due=True 跳过“未到期”校验,避免保护时间反向阻塞本次 poll。 task.next_poll_at = queue_hold_until await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列", detail={ "pipeline_stage": task.pipeline_stage, "payload": redis_payload, "due_next_poll_at": original_next_poll_at, "queue_hold_until": queue_hold_until, }, ) poll_generation_task.apply_async( args=[task.id], kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, queue=POLL_QUEUE, countdown=0, ) await register_poll_active( task, check_at=task.next_poll_at, next_poll_at=task.next_poll_at, reason=f"{source}_has_provider_task_id", ) return "recover_poll_has_provider_id" # 未过 deadline,且没有结果 URL / 供应商任务 ID: # 图片同步任务会重新进入 submit_image_task;视频/其它任务会重新创建供应商任务。 # 这里不能投 poll,因为没有 provider_task_id/seedance_task_id 可查询。 recoverable_create_stages = { ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value, } if task.pipeline_stage in recoverable_create_stages: if task.pipeline_stage not in ( ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ): task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value # 刷新更新时间形成创建队列保护窗口,避免 Beat 在任务尚未消费时每轮重复补投。 task.updated_at = current_time # Release the recovery row lock before writing an event through the # independent logging session or talking to the broker. await db.commit() await _remove_poll_active(_chat_registry_id(task)) await log_task_event( task, event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) chatapi_create_generation_task.apply_async( args=[task.id], kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, queue=CeleryQueue.GEN_CHATAPI_CREATE.value, countdown=0, task_id=( f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:" f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}" ), ) return "recover_create_no_remote_no_provider_before_deadline" # result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。 if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value task.updated_at = current_time await db.commit() await _remove_poll_active(_chat_registry_id(task)) await log_task_event( task, event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) chatapi_create_generation_task.apply_async( args=[task.id], kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, queue=CeleryQueue.GEN_CHATAPI_CREATE.value, countdown=0, task_id=( f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:" f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}" ), ) return "recover_create_result_ready_no_url_before_deadline" return f"skip_stage_{task.pipeline_stage}" async def recover_image_main_create_tasks_once(db: AsyncSession) -> dict[str, int]: """批量恢复同步多图主任务;锁活跃或 DB lease 未过期时绝不接管。""" from app.tasks.generation_create_tasks import chatapi_create_generation_task results: dict[str, int] = {} cursor: str | None = None batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 100)) allowed_stages = [ ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ChatGenerationPipelineStage.PROVIDER_RESULT_STAGED.value, ] while True: query = ( select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value, ChatGenerationTask.gen_type == GenerationType.IMAGE.value, ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.pipeline_stage.in_(allowed_stages), or_( ChatGenerationTask.provider_create_lease_until <= _now(), ( ChatGenerationTask.provider_create_lease_until.is_(None) & ( ChatGenerationTask.updated_at <= _now() - timedelta( seconds=max( 1, int(settings.GENERATION_CREATE_QUEUE_TIMEOUT_SECONDS or 300), ) ) ) ), ), ) .order_by(ChatGenerationTask.id.asc()) .limit(batch_size) .with_for_update(skip_locked=True) ) if cursor: query = query.where(ChatGenerationTask.id > cursor) row_result = await db.execute(query) mains = list(row_result.scalars().all()) if not mains: break main_ids = [str(main.id) for main in mains] cursor = main_ids[-1] child_rows = await db.execute( select(ChatGenerationTask.parent_task_id) .where( ChatGenerationTask.parent_task_id.in_(main_ids), ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value, ChatGenerationTask.deleted_at.is_(None), ) .distinct() ) split_parent_ids = {str(value) for value in child_rows.scalars().all() if value} lock_key_by_id = {str(main.id): _create_lock_key(main) for main in mains} lock_values = await runtime_lock_values(list(lock_key_by_id.values())) now = _now() dispatches: list[tuple[str, int, bool]] = [] expired_claim_logs: list[tuple[str, int, str]] = [] for main in mains: main_id = str(main.id) attempt_no = int(main.generation_attempt_no or 1) if lock_values.get(lock_key_by_id[main_id]): results["live_lock"] = results.get("live_lock", 0) + 1 continue if main_id in split_parent_ids: main.provider_create_claim_token = None main.provider_create_lease_until = None results["already_split"] = results.get("already_split", 0) + 1 continue lease_until = ensure_aware_utc(main.provider_create_lease_until) if main.provider_create_claim_token and lease_until and lease_until > now: results["waiting_db_lease"] = results.get("waiting_db_lease", 0) + 1 continue has_staged_result = bool( main.pipeline_stage == ChatGenerationPipelineStage.PROVIDER_RESULT_STAGED.value and main.provider_response_json ) if _is_expired(main.deadline_at, now) and not has_staged_result: main.provider_create_claim_token = None main.provider_create_lease_until = None await mark_chat_generation_task_failed_and_refund_once( db, task=main, error_message="图片生成任务超时,系统已自动退回本轮媒体生成积分", pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) results["timeout"] = results.get("timeout", 0) + 1 continue old_claim = str(main.provider_create_claim_token or "") main.provider_create_claim_token = None main.provider_create_lease_until = None if not has_staged_result: main.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value if old_claim: expired_claim_logs.append((main_id, attempt_no, str(main.generation_mode))) dispatches.append((main_id, attempt_no, has_staged_result)) await db.commit() for main_id, attempt_no, generation_mode in expired_claim_logs: await log_task_event( task_id=main_id, generation_attempt_no=attempt_no, generation_mode=generation_mode, event_type=ChatGenerationTaskEventType.IMAGE_MAIN_CLAIM_EXPIRED.value, message="图片主任务执行锁已失效且数据库租约已过期,恢复重新投递", ) for main_id, attempt_no, has_staged_result in dispatches: try: chatapi_create_generation_task.apply_async( args=[main_id], kwargs={ "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": attempt_no, }, queue=CeleryQueue.GEN_CHATAPI_CREATE.value, countdown=0, task_id=( f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:" f"{main_id}:attempt:{attempt_no}" ), ) key = "recover_staged_split" if has_staged_result else "recover_create" results[key] = results.get(key, 0) + 1 except Exception: logger.exception("恢复投递图片主任务失败 task_id=%s", main_id) results["enqueue_failed"] = results.get("enqueue_failed", 0) + 1 if len(mains) < batch_size: break return results async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: """启动时生成链路容灾扫描。 启动容灾由 worker_ready 触发,只跑一次完整恢复;周期性视频到期轮询由 Celery Beat 调度 dispatch_due_poll_tasks。 恢复顺序: 1. Redis poll active_index 到期任务; 2. DB fallback 扫描 queued/creating/waiting_remote/polling/result_ready; 3. 下载阶段仍由 recover_download_tasks_once 兜底。 """ checked_ids: set[str] = set() results: dict[str, int] = {} image_main_results = await recover_image_main_create_tasks_once(db) for key, value in image_main_results.items(): results[f"image_main_{key}"] = value due_poll_ids = await redis_get_due_registry_ids( zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, limit=int(settings.POLL_RECOVERY_BATCH_SIZE or settings.GENERATION_RECOVERY_BATCH_SIZE or 100), log_context="poll_active", ) poll_payloads = await redis_get_registry_payloads( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, item_ids=due_poll_ids, log_context="poll_active", ) poll_refs = { registry_item_id: parse_redis_owner_item_id(registry_item_id) for registry_item_id in due_poll_ids } for registry_item_id, ref in poll_refs.items(): if ref.owner_type != GenerationOwnerType.CHAT_GENERATION_TASK.value: continue task = await _load_chat_task_for_update(db, ref.owner_id) if task is None: await db.rollback() await _remove_poll_active(registry_item_id) action = "clean_missing_poll_task" elif ( ref.generation_attempt_no is not None and int(task.generation_attempt_no or 1) != int(ref.generation_attempt_no) ): await db.rollback() await _remove_poll_active(registry_item_id) action = "clean_stale_poll_attempt" else: current_registry_id = _chat_registry_id(task) if registry_item_id != current_registry_id: await _remove_poll_active(registry_item_id) checked_ids.add(task.id) action = await recover_one_generation_task( db, task, payload=poll_payloads.get(registry_item_id), source="startup_poll_redis", ) results[action] = results.get(action, 0) + 1 batch_size = int(settings.GENERATION_RECOVERY_BATCH_SIZE or settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100) max_rounds = max(1, int(settings.GENERATION_RECOVERY_MAX_ROUNDS or 1)) total_db_checked = 0 for _round in range(max_rounds): query_result = await db.execute( select(ChatGenerationTask.id) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.pipeline_stage.in_( [ "queued", "preparing", "creating_provider_task", "waiting_remote", "polling", "result_ready", ] ), ) .order_by(ChatGenerationTask.updated_at.asc()) .limit(batch_size) ) task_ids = [str(value) for value in query_result.scalars().all()] if not task_ids: break progressed_this_round = 0 for task_id in task_ids: if task_id in checked_ids: continue task = await _load_chat_task_for_update(db, task_id) if task is None: await db.rollback() continue action = await recover_one_generation_task( db, task, payload=None, source="startup_db", ) results[action] = results.get(action, 0) + 1 checked_ids.add(task_id) total_db_checked += 1 progressed_this_round += 1 if len(task_ids) < batch_size or progressed_this_round <= 0: break # 子任务可能在 worker 中断前已进入终态但主任务尚未汇总,按稳定游标完整重算全部主任务。 from app.services.generation.ai.task_group_service import aggregate_main_tasks_status_batch reconciled = 0 main_cursor: str | None = None while True: main_query = select(ChatGenerationTask.id).where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value, ) if main_cursor: main_query = main_query.where(ChatGenerationTask.id > main_cursor) main_result = await db.execute(main_query.order_by(ChatGenerationTask.id.asc()).limit(batch_size)) parent_ids = list(main_result.scalars().all()) if not parent_ids: break main_cursor = str(parent_ids[-1]) await aggregate_main_tasks_status_batch(db, parent_task_ids=[str(value) for value in parent_ids]) await db.commit() reconciled += len(parent_ids) if len(parent_ids) < batch_size: break if reconciled: results["reconcile_main"] = reconciled return { "checked": len(checked_ids), "db_checked": total_db_checked, "results": results, } async def recover_stale_create_tasks_once(db: AsyncSession) -> dict[str, Any]: """轻量恢复长时间未消费的创建阶段 ChatGenerationTask。 只扫描 queued/preparing/creating_provider_task,避免周期任务重复执行完整 provider poll、下载和主任务汇总逻辑。 """ current_time = _now() cutoff = current_time - timedelta( seconds=max(1, int(settings.GENERATION_CREATE_QUEUE_TIMEOUT_SECONDS or 300)) ) lease_expired_at = current_time batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20)) result = await db.execute( select(ChatGenerationTask.id) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.pipeline_stage.in_( [ ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ] ), ChatGenerationTask.remote_result_url.is_(None), ChatGenerationTask.provider_task_id.is_(None), ChatGenerationTask.seedance_task_id.is_(None), ChatGenerationTask.updated_at <= cutoff, ( ChatGenerationTask.provider_create_lease_until.is_(None) | (ChatGenerationTask.provider_create_lease_until <= lease_expired_at) ), ) .order_by(ChatGenerationTask.updated_at.asc(), ChatGenerationTask.id.asc()) .limit(batch_size) ) task_ids = [str(value) for value in result.scalars().all()] counts: dict[str, int] = {} image_main_counts = await recover_image_main_create_tasks_once(db) for key, value in image_main_counts.items(): counts[f"image_main_{key}"] = value for task_id in task_ids: task = await _load_chat_task_for_update(db, task_id) if task is None: await db.rollback() action = "skip_missing_task" else: action = await recover_one_generation_task( db, task, payload=None, source="periodic_create_recovery", ) counts[action] = counts.get(action, 0) + 1 return {"checked": len(task_ids), "results": counts} async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]: """周期性轻量到期轮询调度。 只处理视频任务的 next_poll_at 到期记录,不替代启动容灾 recover_generation_tasks_once。 Beat 每分钟触发本任务,本任务把真实供应商轮询投递到 gen_provider_poll 队列。 """ from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active current_time = _now() batch_size = max(1, int(settings.POLL_DUE_DISPATCH_BATCH_SIZE or 100)) poll_lease_expired_at = current_time - timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)) await apply_short_lock_timeout(db) query_result = await db.execute( select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.gen_type == GenerationType.VIDEO.value, ChatGenerationTask.next_poll_at.is_not(None), ChatGenerationTask.next_poll_at <= current_time, ChatGenerationTask.pipeline_stage.in_( [ ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value, ] ), ) .order_by(ChatGenerationTask.next_poll_at.asc(), ChatGenerationTask.updated_at.asc()) .limit(batch_size) .with_for_update(skip_locked=True) ) tasks = query_result.scalars().all() results: dict[str, int] = {} dispatched_task_ids: list[str] = [] dispatched_due_next_poll_at_by_id: dict[str, datetime | None] = {} dispatched_queue_hold_until_by_id: dict[str, datetime] = {} post_commit_logs: list[dict[str, Any]] = [] for task in tasks: action = "skip_unknown" try: if task.pipeline_stage == ChatGenerationPipelineStage.POLLING.value: last_poll_at = ensure_aware_utc(task.last_poll_at) if last_poll_at and last_poll_at > poll_lease_expired_at: await register_poll_active( task, check_at=_poll_queue_timeout_at(last_poll_at), next_poll_at=task.next_poll_at, reason="due_dispatch_polling_lease_alive", ) action = "skip_polling_lease_alive" continue if not (task.seedance_task_id or task.provider_task_id): # Defer FK-backed event logging until the batch row locks are released. post_commit_logs.append({ "task_id": str(task.id), "generation_attempt_no": int(task.generation_attempt_no or 1), "event_type": ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value, "message": "视频到期轮询调度跳过:缺少外部任务ID", "detail": {"pipeline_stage": task.pipeline_stage, "next_poll_at": task.next_poll_at}, }) action = "skip_no_provider_task_id" continue original_next_poll_at = ensure_aware_utc(task.next_poll_at) queue_hold_until = _poll_queue_timeout_at(current_time) task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value # 设置一个队列消费保护时间,避免 Beat 下一分钟看到旧 next_poll_at 又重复投递。 # poll worker 会通过 force_due=True 消费本次到期任务,避免该保护时间被误判为业务未到期。 task.next_poll_at = queue_hold_until dispatched_due_next_poll_at_by_id[task.id] = original_next_poll_at dispatched_queue_hold_until_by_id[task.id] = queue_hold_until dispatched_task_ids.append(task.id) action = "dispatch_poll" except Exception as exc: logger.exception("视频到期轮询调度单条处理失败。task_id=%s", getattr(task, "id", None)) action = "error" post_commit_logs.append({ "task_id": str(getattr(task, "id", "") or ""), "generation_attempt_no": int(getattr(task, "generation_attempt_no", 1) or 1), "event_type": ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value, "message": f"视频到期轮询调度单条处理失败:{exc}", "detail": None, }) finally: results[action] = results.get(action, 0) + 1 await db.commit() for item in post_commit_logs: if item["task_id"]: await log_task_event( task_id=item["task_id"], generation_attempt_no=item["generation_attempt_no"], event_type=item["event_type"], message=item["message"], detail=item["detail"], ) fresh_tasks = [] if dispatched_task_ids: fresh_result = await db.execute( select(ChatGenerationTask) .where( ChatGenerationTask.id.in_(dispatched_task_ids), ChatGenerationTask.deleted_at.is_(None), ) .execution_options(populate_existing=True) ) fresh_tasks = fresh_result.scalars().all() enqueued_count = 0 for task in fresh_tasks: queue_hold_until = dispatched_queue_hold_until_by_id.get(task.id) or ensure_aware_utc(task.next_poll_at) or _poll_queue_timeout_at(current_time) due_next_poll_at = dispatched_due_next_poll_at_by_id.get(task.id) await register_poll_active( task, check_at=queue_hold_until, next_poll_at=queue_hold_until, reason="due_dispatch_poll_queued", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.POLL_DISPATCH_DUE.value, message="视频 next_poll_at 到期,已投递 provider poll 队列", detail={ "due_next_poll_at": due_next_poll_at, "queue_hold_until": queue_hold_until, "queue": POLL_QUEUE, }, ) poll_generation_task.apply_async( args=[task.id], kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, queue=POLL_QUEUE, countdown=0, ) enqueued_count += 1 # 如果 log_task_event 内部不 commit,这里要提交一次 if fresh_tasks: await db.commit() return { "checked": len(tasks), "dispatched": len(dispatched_task_ids), "enqueued": enqueued_count, "results": results, }