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()