from __future__ import annotations import logging import asyncio import errno import math import uuid from datetime import datetime, timedelta, timezone from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.enums.celery_queue import CeleryQueue, CeleryTaskName from app.enums.celery_runtime import CeleryRuntimeDomain from app.enums.generation_status import GenerationRecordPipelineStage from app.enums.generation_task import ( ALLOWED_GENERATION_MODES, ChatGenerationPipelineStage, ChatGenerationTaskEventType, GenerationMode, GenerationOwnerType, GenerationType, ) from app.models.base import async_session from app.models.chat_generation_task import ChatGenerationTask from app.services.celery_download_recovery_service import ( build_download_active_payload, remove_download_active, upsert_download_active, ) from app.services.error_codes import extract_error_message from app.services.generation.download_service import ( download_generation_result, download_video_upscale_source, ) from app.services.generation.log_service import log_task_event from app.services.generation.pipeline.db_lock_service import DatabaseRowLockBusy from app.services.generation.pipeline.lifecycle_service import ( mark_owner_failed_and_refund_once, notify_owner_finished, ) from app.services.generation.pipeline.owner_service import ( GenerationOwner, is_attempt_current, load_generation_owner, load_generation_owner_for_update_retry, normalize_owner_type, owner_is_completed, owner_is_generating, owner_type_of, redis_owner_item_id, set_owner_completed, renew_generation_owner_claim_lease, ) from app.services.media_token_usage_snapshot_service import ( sync_chat_generation_task_media_token_snapshot, sync_generation_record_media_token_snapshot, ) from app.services.resource_accounting_service import ( record_chat_task_generated_resource, record_generation_record_generated_resource, ) from app.services.redis_registry_service import ( RedisExecutionLockError, ensure_aware_utc, ) from app.services.video_upscale.media_service import probe_video from app.services.celery_runtime.runtime_service import CeleryRuntimeLease, RuntimeIdentity from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app logger = logging.getLogger("video_gen") DOWNLOAD_QUEUE = CeleryQueue.GEN_RESULT_DOWNLOAD.value DOWNLOAD_STAGE_QUEUED = ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value DOWNLOAD_STAGE_DOWNLOADING = ChatGenerationPipelineStage.DOWNLOADING.value DOWNLOAD_STAGE_RETRY_WAITING = ChatGenerationPipelineStage.RETRY_WAITING.value def _now() -> datetime: return datetime.now(timezone.utc) def _stage(owner: GenerationOwner, chat_stage: ChatGenerationPipelineStage) -> str: if isinstance(owner, ChatGenerationTask): return chat_stage.value try: return GenerationRecordPipelineStage(chat_stage.value).value except ValueError: return chat_stage.value def _registry_id(owner: GenerationOwner) -> str: return redis_owner_item_id( owner_type_of(owner), owner.id, int(owner.generation_attempt_no or 1) ) def _lock_key(owner_type: str, owner_id: str, attempt_no: int) -> str: return ( f"{settings.GENERATION_DOWNLOAD_LOCK_KEY_PREFIX}:" f"{owner_type}:{owner_id}:attempt:{attempt_no}" ) def _queue_timeout_at() -> datetime: return _now() + timedelta( seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300) ) def _lease_until() -> datetime: return _now() + timedelta( seconds=int(settings.DOWNLOAD_TASK_LEASE_SECONDS or 600) ) def _retry_at(attempt: int) -> datetime: return _now() + timedelta( seconds=max( 1, int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30) * max(1, attempt), ) ) def _countdown_until(value: datetime | None) -> int: target = ensure_aware_utc(value) if target is None: return 1 return max( 1, math.ceil((target - _now()).total_seconds()) + int(settings.DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS or 1), ) def _build_celery_task_id(owner: GenerationOwner, *, reason: str) -> str: return ( f"download:{owner_type_of(owner)}:{owner.id}:" f"attempt:{int(owner.generation_attempt_no or 1)}:" f"{reason[:24]}:{uuid.uuid4().hex[:10]}" ) def _is_non_retryable_download_error(exc: Exception) -> bool: if not bool(getattr(settings, "DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS", True)): return False if isinstance(exc, PermissionError): return True if isinstance(exc, OSError) and getattr(exc, "errno", None) in { errno.EACCES, errno.EPERM, errno.ENOSPC, errno.EROFS, errno.ENAMETOOLONG, }: return True message = str(exc).lower() return any( value in message for value in ( "permission denied", "no space left on device", "read-only file system", "file name too long", ) ) async def _register_active( owner: GenerationOwner, *, check_at: datetime, priority: int, reason: str, ) -> None: item_id = _registry_id(owner) payload = build_download_active_payload( record_id=item_id, celery_task_id=owner.download_celery_task_id, stage=owner.pipeline_stage or "", attempt=owner.download_attempt_count or 0, queue=DOWNLOAD_QUEUE, priority=priority, enqueue_at=owner.download_enqueued_at, started_at=owner.download_started_at, lease_until=owner.download_lease_until, next_retry_at=owner.download_next_retry_at, check_at=check_at, reason=reason, ) payload.update( { "owner_type": owner_type_of(owner), "owner_id": owner.id, "generation_attempt_no": int(owner.generation_attempt_no or 1), "download_claim_token_suffix": ( str(owner.download_claim_token)[-8:] if owner.download_claim_token else None ), } ) await upsert_download_active( record_id=item_id, payload=payload, check_at=check_at ) async def _remove_active(owner: GenerationOwner) -> None: await remove_download_active(_registry_id(owner)) async def _apply( owner: GenerationOwner, *, priority: int, countdown: int | None, reason: str, ) -> None: download_generation_result_task.apply_async( args=[owner.id], kwargs={ "owner_type": owner_type_of(owner), "generation_attempt_no": int(owner.generation_attempt_no or 1), }, queue=DOWNLOAD_QUEUE, priority=priority, countdown=countdown, task_id=owner.download_celery_task_id, ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE.value, to_stage=owner.pipeline_stage, detail={ "reason": reason, "celery_task_id": owner.download_celery_task_id, }, ) def _recover_enqueue_due(owner: GenerationOwner, now: datetime) -> bool: result_ready = _stage(owner, ChatGenerationPipelineStage.RESULT_READY) queued = _stage(owner, ChatGenerationPipelineStage.DOWNLOAD_QUEUED) downloading = _stage(owner, ChatGenerationPipelineStage.DOWNLOADING) retry_waiting = _stage(owner, ChatGenerationPipelineStage.RETRY_WAITING) if owner.pipeline_stage == result_ready: return True if owner.pipeline_stage == queued: enqueued_at = ensure_aware_utc(owner.download_enqueued_at) if enqueued_at is None: return True return enqueued_at + timedelta( seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300) ) <= now if owner.pipeline_stage == downloading: lease_until = ensure_aware_utc(owner.download_lease_until) return lease_until is None or lease_until <= now if owner.pipeline_stage == retry_waiting: next_retry_at = ensure_aware_utc(owner.download_next_retry_at) return next_retry_at is None or next_retry_at <= now return False async def enqueue_download_task( db: AsyncSession, owner: GenerationOwner, *, recover: bool = False, reason: str | None = None, countdown: int | None = None, ) -> str | None: reason = reason or ("recover" if recover else "normal") if ( isinstance(owner, ChatGenerationTask) and owner.generation_mode not in ALLOWED_GENERATION_MODES ): return None if not owner_is_generating(owner) or not owner.remote_result_url: return None if recover and not _recover_enqueue_due(owner, _now()): return owner.download_celery_task_id owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.DOWNLOAD_QUEUED ) owner.download_celery_task_id = _build_celery_task_id(owner, reason=reason) owner.download_enqueued_at = _now() owner.download_started_at = None owner.download_claim_token = None owner.download_lease_until = None owner.download_next_retry_at = None if not owner.download_storage_date_dir: created_at = ensure_aware_utc(owner.created_at) or _now() owner.download_storage_date_dir = created_at.strftime("%Y/%m/%d") await db.commit() priority = int( settings.DOWNLOAD_TASK_PRIORITY_RECOVER if recover else settings.DOWNLOAD_TASK_PRIORITY_NORMAL ) await _register_active( owner, check_at=_queue_timeout_at(), priority=priority, reason=reason, ) try: await _apply( owner, priority=priority, countdown=countdown, reason=reason, ) except Exception as exc: await log_task_event( owner, event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED.value, message="下载任务投递失败,等待恢复扫描", detail={"reason": reason, "error": str(exc)}, ) return None return owner.download_celery_task_id async def _claim( db: AsyncSession, owner: GenerationOwner, *, claim_token: str, ) -> bool: if not owner_is_generating(owner) or owner_is_completed(owner): return False allowed = { _stage(owner, ChatGenerationPipelineStage.RESULT_READY), _stage(owner, ChatGenerationPipelineStage.DOWNLOAD_QUEUED), _stage(owner, ChatGenerationPipelineStage.DOWNLOADING), _stage(owner, ChatGenerationPipelineStage.RETRY_WAITING), } if owner.pipeline_stage not in allowed: return False now = _now() # Redis execution lock is authoritative. A database lease left by a # crashed worker is audit/recovery metadata and must not block the worker # that owns the current Redis download lock. next_retry = ensure_aware_utc(owner.download_next_retry_at) if ( owner.pipeline_stage == _stage(owner, ChatGenerationPipelineStage.RETRY_WAITING) and next_retry and next_retry > now ): return False owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.DOWNLOADING ) owner.download_claim_token = claim_token owner.download_started_at = now owner.download_lease_until = _lease_until() owner.download_attempt_count = int(owner.download_attempt_count or 0) + 1 owner.download_last_error = None await db.commit() await _register_active( owner, check_at=owner.download_lease_until, priority=int(settings.DOWNLOAD_TASK_PRIORITY_NORMAL), reason="downloading_lease", ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.DOWNLOAD_START.value, to_stage=owner.pipeline_stage, detail={ "download_attempt_count": owner.download_attempt_count, "claim_token_suffix": claim_token[-8:], }, ) return True async def _sync_snapshot(db: AsyncSession, owner: GenerationOwner) -> None: if isinstance(owner, ChatGenerationTask): await sync_chat_generation_task_media_token_snapshot(db, owner) else: await sync_generation_record_media_token_snapshot( db, owner, provider_response=owner.provider_response_json ) async def _record_resource( db: AsyncSession, owner: GenerationOwner, downloaded ) -> None: kwargs = dict( resource_url=downloaded.url, storage_path=downloaded.storage_path, file_size_bytes=downloaded.file_size_bytes, remote_url=owner.remote_result_url, generated_at=owner.generated_at, ) if isinstance(owner, ChatGenerationTask): await record_chat_task_generated_resource(db, owner, **kwargs) else: await record_generation_record_generated_resource(db, owner, **kwargs) async def _mark_failed( db: AsyncSession, owner: GenerationOwner, exc: Exception, *, non_retryable: bool, ) -> None: error_message = ( extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc) ) image_child = bool( isinstance(owner, ChatGenerationTask) and owner.gen_type == GenerationType.IMAGE.value and owner.generation_mode == GenerationMode.CHATAPI_CHILD.value ) if image_child: owner.status = "failed" owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.DOWNLOAD_FAILED ) owner.error_message = error_message else: await mark_owner_failed_and_refund_once( db, owner, error_message=error_message, pipeline_stage=_stage( owner, ChatGenerationPipelineStage.DOWNLOAD_FAILED ), ) owner.download_claim_token = None owner.download_last_error = error_message owner.download_lease_until = None owner.download_next_retry_at = None await db.commit() await notify_owner_finished(db, owner) await db.commit() await _remove_active(owner) await log_task_event( owner, event_type=( ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value if non_retryable else ChatGenerationTaskEventType.DOWNLOAD_FAILED.value ), message=error_message, ) async def _schedule_retry( db: AsyncSession, owner: GenerationOwner, exc: Exception ) -> None: error_message = ( extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc) ) owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.RETRY_WAITING ) owner.download_next_retry_at = _retry_at( int(owner.download_attempt_count or 1) ) owner.download_claim_token = None owner.download_lease_until = None owner.download_last_error = error_message owner.download_celery_task_id = _build_celery_task_id( owner, reason="retry" ) await db.commit() await _register_active( owner, check_at=owner.download_next_retry_at, priority=int(settings.DOWNLOAD_TASK_PRIORITY_RECOVER), reason="download_retry_waiting", ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING.value, message=error_message, detail={"next_retry_at": owner.download_next_retry_at}, ) await _apply( owner, priority=int(settings.DOWNLOAD_TASK_PRIORITY_RECOVER), countdown=_countdown_until(owner.download_next_retry_at), reason="download_exception_retry", ) async def _resolve_attempt( task_id: str, *, owner_type: str, message_attempt: int | None ) -> int | None: async with async_session() as db: owner = await load_generation_owner( db, owner_type=owner_type, owner_id=task_id, for_update=False ) if not owner: await remove_download_active( redis_owner_item_id(owner_type, task_id, message_attempt) ) return None if not is_attempt_current(owner, message_attempt): await remove_download_active( redis_owner_item_id(owner_type, task_id, message_attempt) ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.STALE_ATTEMPT_MESSAGE_SKIPPED.value, message="下载消息属于旧生成轮次,已跳过", ) return None return int(owner.generation_attempt_no or 1) async def _restore_after_lock_error( *, owner_type: str, owner_id: str, attempt_no: int, claim_token: str, ) -> None: async with async_session() as db: owner = await load_generation_owner( db, owner_type=owner_type, owner_id=owner_id, for_update=True, ) if not owner or not is_attempt_current(owner, attempt_no): return if owner.download_claim_token != claim_token: return owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.DOWNLOAD_QUEUED ) owner.download_claim_token = None owner.download_lease_until = None owner.download_started_at = None owner.download_attempt_count = max( 0, int(owner.download_attempt_count or 0) - 1 ) owner.download_enqueued_at = _now() await db.commit() await _register_active( owner, check_at=_queue_timeout_at(), priority=int(settings.DOWNLOAD_TASK_PRIORITY_RECOVER), reason="download_execution_lock_error", ) async def _run( task_id: str, *, owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value, generation_attempt_no: int | None = None, ): normalized_owner_type = normalize_owner_type(owner_type) effective_attempt = await _resolve_attempt( task_id, owner_type=normalized_owner_type, message_attempt=generation_attempt_no, ) if effective_attempt is None: return token = uuid.uuid4().hex lease = await CeleryRuntimeLease.acquire( identity=RuntimeIdentity( domain=CeleryRuntimeDomain.GENERATION_DOWNLOAD.value, owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, task_name=CeleryTaskName.DOWNLOAD_GENERATION_RESULT.value, queue=DOWNLOAD_QUEUE, registry_item_id=redis_owner_item_id( normalized_owner_type, task_id, effective_attempt ), ), lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt), hash_key=settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY, zset_key=settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY, token=token, ttl_seconds=int(settings.GENERATION_DOWNLOAD_LOCK_TTL_SECONDS or 600), heartbeat_interval_seconds=int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30), pipeline_stage=ChatGenerationPipelineStage.DOWNLOADING.value, db_heartbeat=lambda owned_token: renew_generation_owner_claim_lease( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, claim_field="download_claim_token", lease_field="download_lease_until", token=owned_token, lease_seconds=int(settings.DOWNLOAD_TASK_LEASE_SECONDS or 600), ), ) if lease is None: return claimed = False try: async with async_session() as db: owner = await load_generation_owner( db, owner_type=normalized_owner_type, owner_id=task_id, for_update=True, ) if not owner: await remove_download_active( redis_owner_item_id( normalized_owner_type, task_id, effective_attempt ) ) return if not is_attempt_current(owner, effective_attempt): await _remove_active(owner) return if not await _claim(db, owner, claim_token=lease.token): return claimed = True try: use_upscale = bool( owner.gen_type == GenerationType.VIDEO.value and owner.video_upscale_enabled_snapshot and owner.video_upscale_snapshot_json ) try: async with asyncio.timeout( max(1, int(settings.GENERATION_DOWNLOAD_TOTAL_TIMEOUT_SECONDS or 480)) ): downloaded = ( await download_video_upscale_source( owner, execution_guard=lease.ensure_owned ) if use_upscale else await download_generation_result( owner, execution_guard=lease.ensure_owned ) ) except TimeoutError as exc: raise TimeoutError("下载总耗时超过限制") from exc # ffprobe is an external process and must not run while holding # the generation owner row lock. upscale_source_info = None if use_upscale: await lease.ensure_owned() upscale_source_info = await probe_video( str(downloaded.storage_path or "") ) await lease.ensure_owned() owner = await load_generation_owner_for_update_retry( db, owner_type=normalized_owner_type, owner_id=task_id, ) if not owner or not is_attempt_current(owner, effective_attempt): await remove_download_active( redis_owner_item_id( normalized_owner_type, task_id, effective_attempt ) ) return if owner.download_claim_token != lease.token: return if use_upscale: from app.services.video_upscale.task_service import ( enqueue_upscale_task, prepare_video_upscale_task, ) kwargs = ( {"task": owner} if isinstance(owner, ChatGenerationTask) else {"generation_record": owner} ) upscale = await prepare_video_upscale_task( db, **kwargs, source_local_path=str(downloaded.storage_path or ""), source_file_size_bytes=downloaded.file_size_bytes, source_remote_url=owner.remote_result_url, source_info=upscale_source_info, ) owner.download_claim_token = None owner.download_lease_until = None owner.download_next_retry_at = None owner.download_last_error = None await db.commit() await _remove_active(owner) await enqueue_upscale_task( db, upscale=upscale, reason="source_download_completed" ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value, to_stage=owner.pipeline_stage, detail={ "upscale_source_path": downloaded.storage_path }, ) return if owner.gen_type == GenerationType.IMAGE.value: owner.image_url = downloaded.url else: owner.video_url = downloaded.url owner.video_cover_url = downloaded.cover_url set_owner_completed(owner) owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.DONE ) owner.generated_at = _now() owner.retry_count = int( getattr(owner, "manual_retry_count", 0) or 0 ) owner.download_claim_token = None owner.download_lease_until = None owner.download_next_retry_at = None owner.download_last_error = None await _record_resource(db, owner, downloaded) await _sync_snapshot(db, owner) await db.commit() await notify_owner_finished(db, owner) await db.commit() await _remove_active(owner) await log_task_event( owner, event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value, to_stage=owner.pipeline_stage, detail={ "resource_url": downloaded.url, "file_size_bytes": downloaded.file_size_bytes, }, ) except (RedisExecutionLockError, DatabaseRowLockBusy): try: await db.rollback() except Exception: pass raise except Exception as exc: try: await db.rollback() except Exception: pass await lease.ensure_owned() owner = await load_generation_owner_for_update_retry( db, owner_type=normalized_owner_type, owner_id=task_id, ) if not owner or not is_attempt_current(owner, effective_attempt): await remove_download_active( redis_owner_item_id( normalized_owner_type, task_id, effective_attempt ) ) return if owner.download_claim_token != lease.token: return non_retryable = _is_non_retryable_download_error(exc) max_attempts = int(settings.DOWNLOAD_TASK_MAX_ATTEMPTS or 3) if non_retryable or int(owner.download_attempt_count or 0) >= max_attempts: await _mark_failed( db, owner, exc, non_retryable=non_retryable ) else: await _schedule_retry(db, owner, exc) except (RedisExecutionLockError, DatabaseRowLockBusy): if claimed: try: await _restore_after_lock_error( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, claim_token=lease.token, ) except Exception: logger.exception( "下载执行锁异常后的状态恢复失败 owner_type=%s owner_id=%s attempt=%s", normalized_owner_type, task_id, effective_attempt, ) raise finally: await lease.close() if celery_app: @celery_app.task( name="generation.download_generation_result_task", bind=True, max_retries=3, default_retry_delay=30, ) def download_generation_result_task( self, task_id: str, owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value, generation_attempt_no: int | None = None, ): try: return run_async( _run( task_id, owner_type=owner_type, generation_attempt_no=generation_attempt_no, ) ) except Exception as exc: retries = int(getattr(self.request, "retries", 0) or 0) + 1 countdown = int( settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30 ) * max(1, retries) raise self.retry(exc=exc, countdown=countdown) else: class _DisabledTask: def delay(self, *args, **kwargs): raise RuntimeError("Celery is disabled") def apply_async(self, *args, **kwargs): raise RuntimeError("Celery is disabled") download_generation_result_task = _DisabledTask()