Files
video-gen/video-gen-api/app/tasks/generation_download_tasks.py
T

667 lines
26 KiB
Python

from __future__ import annotations
from app.tasks.async_runner import run_async
import errno
import math
import uuid
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationMode,
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,
ensure_aware_utc,
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
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
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app
DOWNLOAD_QUEUE = "gen_result_download"
DOWNLOAD_STAGE_QUEUED = ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value
DOWNLOAD_STAGE_DOWNLOADING = ChatGenerationPipelineStage.DOWNLOADING.value
DOWNLOAD_STAGE_RETRY_WAITING = ChatGenerationPipelineStage.RETRY_WAITING.value
DOWNLOAD_STAGE_DONE = ChatGenerationPipelineStage.DONE.value
DOWNLOAD_STAGE_FAILED = ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
RESULT_READY_STAGE = ChatGenerationPipelineStage.RESULT_READY.value
def _now() -> datetime:
return datetime.now(timezone.utc)
def _task_created_date_dir(task: ChatGenerationTask) -> str:
created_at = ensure_aware_utc(getattr(task, "created_at", None)) or _now()
return created_at.strftime("%Y/%m/%d")
def _queue_timeout_at(now: datetime | None = None) -> datetime:
now = now or _now()
return now + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
def _lease_until(now: datetime | None = None) -> datetime:
now = now or _now()
return now + timedelta(seconds=int(settings.DOWNLOAD_TASK_LEASE_SECONDS or 600))
def _retry_at(attempt: int, now: datetime | None = None) -> datetime:
now = now or _now()
base = int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30)
return now + timedelta(seconds=max(1, base * max(1, attempt)))
def _countdown_until(value: datetime | None, *, minimum: int = 1) -> int:
target = ensure_aware_utc(value)
if target is None:
return minimum
extra = int(getattr(settings, "DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS", 1) or 0)
seconds = (target - _now()).total_seconds()
return max(minimum, math.ceil(seconds) + extra)
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
value = ensure_aware_utc(value)
if value is None:
return True
return value <= (now or _now())
def _is_already_completed(task: ChatGenerationTask) -> bool:
if task.status == ChatGenerationTaskStatus.COMPLETED.value or task.pipeline_stage == DOWNLOAD_STAGE_DONE:
if task.gen_type == GenerationType.IMAGE.value and task.image_url:
return True
if task.gen_type == GenerationType.VIDEO.value and task.video_url:
return True
return False
def _build_celery_task_id(task_id: str, attempt: int | None = None, reason: str | None = None) -> str:
safe_reason = (reason or "download").replace(" ", "_")[:32]
return f"download:{task_id}:{int(attempt or 0)}:{safe_reason}:{uuid.uuid4().hex[:12]}"
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()
non_retryable_fragments = (
"permission denied",
"no space left on device",
"read-only file system",
"file name too long",
"invalid argument",
)
return any(fragment in message for fragment in non_retryable_fragments)
async def _log_download_event(
task: ChatGenerationTask | None = None,
*,
task_id: str | None = None,
event_type: ChatGenerationTaskEventType | str,
from_status: str | None = None,
to_status: str | None = None,
from_stage: str | None = None,
to_stage: str | None = None,
message: str | None = None,
detail: Any = None,
) -> None:
if not bool(getattr(settings, "DOWNLOAD_EVENT_VERBOSE_ENABLED", True)):
# 成功/失败关键事件仍保留;只关闭 verbose skip 事件。
critical = {
ChatGenerationTaskEventType.DOWNLOAD_START.value,
ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
ChatGenerationTaskEventType.DOWNLOAD_FAILED.value,
ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value,
ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING.value,
}
event_value = event_type.value if hasattr(event_type, "value") else str(event_type)
if event_value not in critical:
return
await log_task_event(
task,
task_id=task_id,
event_type=event_type.value if hasattr(event_type, "value") else str(event_type),
from_status=from_status,
to_status=to_status,
from_stage=from_stage,
to_stage=to_stage,
message=message,
detail=detail,
)
async def _register_active_from_task(
task: ChatGenerationTask,
*,
check_at: datetime,
priority: int,
reason: str | None = None,
) -> None:
payload = build_download_active_payload(
record_id=task.id,
celery_task_id=task.download_celery_task_id,
stage=task.pipeline_stage or "",
attempt=task.download_attempt_count or 0,
queue=DOWNLOAD_QUEUE,
priority=priority,
enqueue_at=task.download_enqueued_at,
started_at=task.download_started_at,
lease_until=task.download_lease_until,
next_retry_at=task.download_next_retry_at,
check_at=check_at,
reason=reason,
)
await upsert_download_active(record_id=task.id, payload=payload, check_at=check_at)
async def _apply_download_async(
task: ChatGenerationTask,
*,
priority: int,
countdown: int | None,
reason: str,
event_type: ChatGenerationTaskEventType,
failed_event_type: ChatGenerationTaskEventType,
) -> bool:
if not celery_app:
await _log_download_event(
task,
event_type=failed_event_type,
message="Celery 未启用,下载任务无法投递",
detail={"queue": DOWNLOAD_QUEUE, "priority": priority, "countdown": countdown, "reason": reason},
)
return False
try:
download_generation_result_task.apply_async(
args=[task.id],
queue=DOWNLOAD_QUEUE,
priority=priority,
countdown=countdown,
task_id=task.download_celery_task_id,
)
await _log_download_event(
task,
event_type=event_type,
to_stage=task.pipeline_stage,
detail={
"queue": DOWNLOAD_QUEUE,
"priority": priority,
"countdown": countdown,
"download_celery_task_id": task.download_celery_task_id,
"reason": reason,
"attempt": task.download_attempt_count,
},
)
return True
except Exception as exc:
await _log_download_event(
task,
event_type=failed_event_type,
message=str(exc),
detail={
"queue": DOWNLOAD_QUEUE,
"priority": priority,
"countdown": countdown,
"download_celery_task_id": task.download_celery_task_id,
"reason": reason,
"attempt": task.download_attempt_count,
},
)
return False
async def enqueue_download_task(
db: AsyncSession,
task: ChatGenerationTask,
*,
recover: bool = False,
reason: str | None = None,
countdown: int | None = None,
) -> str | None:
"""统一投递图片/视频下载任务,并同步 DB + Redis active 注册表。"""
reason = reason or ("recover" if recover else "normal")
if not task:
return None
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="不支持的 generation_mode", detail={"reason": reason})
return None
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="任务不是 generating 状态", detail={"reason": reason, "status": task.status})
return None
if not task.remote_result_url:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="缺少 remote_result_url", detail={"reason": reason})
return None
now = _now()
priority = settings.DOWNLOAD_TASK_PRIORITY_RECOVER if recover else settings.DOWNLOAD_TASK_PRIORITY_NORMAL
celery_task_id = _build_celery_task_id(
task.id,
attempt=task.download_attempt_count or task.retry_count or 0,
reason=reason,
)
old_stage = task.pipeline_stage
task.pipeline_stage = DOWNLOAD_STAGE_QUEUED
task.download_celery_task_id = celery_task_id
task.download_enqueued_at = now
task.download_next_retry_at = None
if not task.download_storage_date_dir:
task.download_storage_date_dir = _task_created_date_dir(task)
await db.commit()
check_at = _queue_timeout_at(now)
try:
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
except Exception as exc:
# Redis active 注册表只用于恢复,不应阻止真实 Celery 投递。
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
message=f"下载恢复注册表写入失败: {exc}",
detail={"reason": reason, "celery_task_id": celery_task_id},
)
applied = await _apply_download_async(
task,
priority=priority,
countdown=countdown,
reason=reason,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
)
if not applied:
# apply_async 失败不能伪装成已投递。保留远程结果,进入下载恢复等待。
try:
await remove_download_active(task.id)
except Exception:
pass
refreshed = await _reload_task(db, task.id)
if refreshed and refreshed.status == ChatGenerationTaskStatus.GENERATING.value:
retry_at = now + timedelta(seconds=int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30))
refreshed.pipeline_stage = DOWNLOAD_STAGE_RETRY_WAITING
refreshed.download_next_retry_at = retry_at
refreshed.download_last_error = "Celery 下载任务投递失败,等待恢复重试"
refreshed.download_lease_until = None
await db.commit()
try:
await _register_active_from_task(
refreshed,
check_at=retry_at,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
reason="enqueue_failed_wait_recovery",
)
except Exception:
pass
return None
if old_stage != DOWNLOAD_STAGE_QUEUED:
# 独立记录阶段变化的上下文,便于和真正投递事件对照。
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
from_stage=old_stage,
to_stage=DOWNLOAD_STAGE_QUEUED,
detail={"reason": reason, "recover": recover, "check_at": check_at},
)
return celery_task_id
async def _reload_task(db: AsyncSession, task_id: str) -> ChatGenerationTask | None:
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1)
)
return result.scalar_one_or_none()
async def _reschedule_not_due_retry(task: ChatGenerationTask, *, now: datetime) -> None:
countdown = _countdown_until(task.download_next_retry_at)
await _register_active_from_task(
task,
check_at=ensure_aware_utc(task.download_next_retry_at) or (now + timedelta(seconds=countdown)),
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
reason="retry_waiting_not_due",
)
await _apply_download_async(
task,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=countdown,
reason="retry_waiting_not_due",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
)
async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> bool:
now = _now()
if not task:
return False
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="下载任务跳过:不支持的 generation_mode")
return False
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="下载任务跳过:任务不是 generating 状态", detail={"status": task.status, "stage": task.pipeline_stage})
return False
if _is_already_completed(task):
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_ALREADY_COMPLETED, message="下载任务跳过:任务已完成", detail={"status": task.status, "stage": task.pipeline_stage})
return False
if not task.remote_result_url:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="下载任务跳过:缺少 remote_result_url")
return False
stage = task.pipeline_stage
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if not _is_expired(task.download_lease_until, now):
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE,
message="下载任务跳过:已有 downloading lease 且未过期",
detail={"lease_until": task.download_lease_until, "download_celery_task_id": task.download_celery_task_id},
)
return False
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_STUCK_RECOVER,
message=f"downloading lease 已过期,重新抢占下载。lease_until={task.download_lease_until}",
)
elif stage == DOWNLOAD_STAGE_RETRY_WAITING:
if not _is_expired(task.download_next_retry_at, now):
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_NOT_DUE,
message="下载重试提前触发,尚未到 next_retry_at,已重新投递延后重试",
detail={
"now": now,
"download_next_retry_at": task.download_next_retry_at,
"download_celery_task_id": task.download_celery_task_id,
},
)
await _reschedule_not_due_retry(task, now=now)
return False
elif stage in (DOWNLOAD_STAGE_QUEUED, RESULT_READY_STAGE):
pass
else:
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_STAGE_NOT_ALLOWED,
message="下载任务跳过:当前阶段不允许下载",
detail={"stage": stage, "status": task.status, "download_celery_task_id": task.download_celery_task_id},
)
return False
old_stage = stage
task.pipeline_stage = DOWNLOAD_STAGE_DOWNLOADING
task.download_started_at = now
task.download_lease_until = _lease_until(now)
task.download_next_retry_at = None
task.download_attempt_count = int(task.download_attempt_count or 0) + 1
task.retry_count = task.download_attempt_count
if not task.download_storage_date_dir:
task.download_storage_date_dir = _task_created_date_dir(task)
await db.commit()
await _register_active_from_task(
task,
check_at=task.download_lease_until,
priority=settings.DOWNLOAD_TASK_PRIORITY_NORMAL,
reason="claim_download_lease",
)
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_START,
from_stage=old_stage,
to_stage=DOWNLOAD_STAGE_DOWNLOADING,
detail={
"attempt": task.download_attempt_count,
"lease_until": task.download_lease_until,
"download_celery_task_id": task.download_celery_task_id,
},
)
return True
async def _mark_retry_waiting(db: AsyncSession, task: ChatGenerationTask, exc: Exception) -> datetime:
now = _now()
attempt = int(task.download_attempt_count or task.retry_count or 0)
next_retry_at = _retry_at(attempt, now)
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
retry_celery_task_id = _build_celery_task_id(task.id, attempt=attempt, reason="retry_waiting")
task.pipeline_stage = DOWNLOAD_STAGE_RETRY_WAITING
task.download_celery_task_id = retry_celery_task_id
task.download_next_retry_at = next_retry_at
task.download_lease_until = None
task.download_last_error = error_message
task.retry_count = attempt
await db.commit()
await _register_active_from_task(
task,
check_at=next_retry_at,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
reason="download_retry_waiting",
)
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING,
message=error_message,
to_stage=DOWNLOAD_STAGE_RETRY_WAITING,
detail={
"attempt": attempt,
"next_retry_at": next_retry_at,
"download_celery_task_id": retry_celery_task_id,
},
)
return next_retry_at
def _should_final_fail(task: ChatGenerationTask) -> bool:
return int(task.download_attempt_count or task.retry_count or 0) >= int(settings.DOWNLOAD_TASK_MAX_ATTEMPTS or 3)
async def _mark_download_failed(
db: AsyncSession,
task: ChatGenerationTask,
*,
exc: Exception,
non_retryable: bool = False,
) -> None:
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
is_image_child = (
task.gen_type == GenerationType.IMAGE.value
and task.generation_mode == GenerationMode.CHATAPI_CHILD.value
)
if is_image_child:
# 图片生成费用属于 main;child 下载失败只记录下载终态,不退图片生成积分。
task.status = ChatGenerationTaskStatus.FAILED.value
task.pipeline_stage = DOWNLOAD_STAGE_FAILED
task.error_message = error_message
await db.flush()
else:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=DOWNLOAD_STAGE_FAILED,
)
task.download_last_error = error_message
task.download_lease_until = None
task.download_next_retry_at = None
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit()
await remove_download_active(task.id)
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE if non_retryable else ChatGenerationTaskEventType.DOWNLOAD_FAILED,
message=task.error_message,
detail={
"download_attempt_count": task.download_attempt_count,
"max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS,
"non_retryable": non_retryable,
"download_last_error": task.download_last_error,
},
)
async def _run(task_id: str):
async with async_session() as db:
task = await _reload_task(db, task_id)
if not task:
await _log_download_event(
task_id=task_id,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING,
message="下载任务跳过:ChatGenerationTask 不存在或已软删",
)
await remove_download_active(task_id)
return
claimed = await _claim_download_lease(db, task)
if not claimed:
return
try:
downloaded = await download_generation_result(task)
task = await _reload_task(db, task_id)
if not task:
await _log_download_event(
task_id=task_id,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING,
message="下载完成后任务不存在或已软删",
)
await remove_download_active(task_id)
return
if task.gen_type == GenerationType.IMAGE.value:
task.image_url = downloaded.url
else:
task.video_url = downloaded.url
task.video_cover_url = downloaded.cover_url
task.status = ChatGenerationTaskStatus.COMPLETED.value
task.pipeline_stage = DOWNLOAD_STAGE_DONE
task.generated_at = _now()
task.retry_count = 0
task.download_lease_until = None
task.download_next_retry_at = None
task.download_last_error = None
await record_chat_task_generated_resource(
db,
task,
resource_url=downloaded.url,
storage_path=downloaded.storage_path,
file_size_bytes=downloaded.file_size_bytes,
remote_url=task.remote_result_url,
generated_at=task.generated_at,
)
await sync_chat_generation_task_media_token_snapshot(db, task)
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit()
await remove_download_active(task.id)
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS,
to_status=ChatGenerationTaskStatus.COMPLETED.value,
to_stage=DOWNLOAD_STAGE_DONE,
detail={
"resource_url": downloaded.url,
"video_cover_url": downloaded.cover_url,
"file_size_bytes": downloaded.file_size_bytes,
"download_attempt_count": task.download_attempt_count,
},
)
except Exception as exc:
try:
await db.rollback()
except Exception:
pass
task = await _reload_task(db, task_id)
if not task:
await remove_download_active(task_id)
return
non_retryable = _is_non_retryable_download_error(exc)
if non_retryable or _should_final_fail(task):
await _mark_download_failed(db, task, exc=exc, non_retryable=non_retryable)
else:
next_retry_at = await _mark_retry_waiting(db, task, exc)
delay_seconds = _countdown_until(next_retry_at)
await _apply_download_async(
task,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=delay_seconds,
reason="download_exception_retry",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
)
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):
try:
return run_async(_run(task_id))
except Exception as exc:
# 只重试 run_async/连接池/worker 中断等基础设施异常;下载业务异常已在 _run 内写入 retry_waiting。
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()