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

201 lines
7.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from app.tasks.async_runner import run_async
from datetime import datetime, timezone, timedelta
from sqlalchemy import select
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
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.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app
# downloading 卡住多久后允许自动恢复。
# 说明:
# - worker 在 pipeline_stage 改成 downloading 后,如果被 kill,任务可能永远停在 downloading。
# - 这里允许超过该时间的 downloading 任务重新进入下载流程。
# - 如果你的视频文件特别大,可以把这个时间调大,比如 20 * 60。
DOWNLOAD_STUCK_SECONDS = 10 * 60
def _to_aware_utc(dt):
"""
把 datetime 统一转成 timezone-aware UTC,避免 offset-naive 和 offset-aware 比较报错。
PostgreSQL / SQLite / 不同驱动返回的 updated_at 可能有时区,也可能没有。
"""
if not dt:
return None
if dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc)
return dt.astimezone(timezone.utc)
def _is_recent_downloading(task: ChatGenerationTask) -> bool:
"""
判断 downloading 是否仍然是较新的下载任务。
返回 True
- 说明可能有另一个 worker 刚进入下载,不要重复下载。
返回 False
- 说明 downloading 已经超过 DOWNLOAD_STUCK_SECONDS,认为可能卡死,可以恢复。
"""
updated_at = _to_aware_utc(getattr(task, "updated_at", None))
if not updated_at:
return False
return datetime.now(timezone.utc) - updated_at < timedelta(seconds=DOWNLOAD_STUCK_SECONDS)
async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
"""
rollback 后重新查询任务对象。
说明:
- SQLAlchemy rollback 后,当前 ORM 对象可能过期。
- 继续访问旧 task 有概率触发异步懒加载异常。
"""
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
)
return result.scalar_one_or_none()
async def _run(task_id: str):
async with async_session() as db:
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
)
task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async":
return
if task.status != "generating":
return
# 关键修改 3
# 原来只允许 result_ready 进入下载。
# 现在允许 downloading 恢复,但只有“卡住超过 DOWNLOAD_STUCK_SECONDS”的 downloading 才继续。
if task.pipeline_stage == "downloading":
if _is_recent_downloading(task):
# downloading 很新,说明可能有 worker 正在下载,直接跳过,避免并发重复下载。
return
# downloading 已经很久没更新,认为 worker 可能挂了,允许恢复下载。
await log_task_event(
task,
event_type="DOWNLOAD_STUCK_RECOVER",
message=f"downloading 超过 {DOWNLOAD_STUCK_SECONDS} 秒,重新进入下载流程",
)
elif task.pipeline_stage != "result_ready":
return
try:
old_stage = task.pipeline_stage
# 无论从 result_ready 进入,还是从 stuck downloading 恢复,都重新标记为 downloading。
task.pipeline_stage = "downloading"
await db.commit()
await log_task_event(
task,
event_type="DOWNLOAD_START",
from_stage=old_stage,
to_stage="downloading",
)
downloaded = await download_generation_result(task)
if task.gen_type == "image":
task.image_url = downloaded.url
else:
task.video_url = downloaded.url
task.status = "completed"
task.pipeline_stage = "done"
task.generated_at = datetime.now(timezone.utc)
task.retry_count = 0
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 db.commit()
await log_task_event(
task,
event_type="DOWNLOAD_SUCCESS",
to_status="completed",
to_stage="done",
detail={
"resource_url": downloaded.url,
"file_size_bytes": downloaded.file_size_bytes,
},
)
except Exception as exc:
# 关键修改 2
# 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。
try:
await db.rollback()
except Exception:
pass
task = await _reload_task(db, task_id)
if not task:
return
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > 3:
task.status = "failed"
task.pipeline_stage = "download_failed"
task.error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
await db.commit()
await log_task_event(
task,
event_type="DOWNLOAD_FAILED",
message=task.error_message,
)
else:
# 下载失败但未超过重试次数,改回 result_ready,等待下一次下载。
# 这样不会卡死在 downloading。
task.pipeline_stage = "result_ready"
await db.commit()
download_generation_result_task.apply_async(
args=[task.id],
countdown=30 * task.retry_count,
)
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):
return run_async(_run(task_id))
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()