203 lines
7.1 KiB
Python
203 lines
7.1 KiB
Python
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.video_cover_url = downloaded.cover_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,
|
||
"video_cover_url": downloaded.cover_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()
|