379 lines
13 KiB
Python
379 lines
13 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
from datetime import datetime, timedelta, timezone
|
|
from typing import Any, Dict
|
|
|
|
from sqlalchemy import or_, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.config import settings
|
|
from app.models.chat_generation_task import ChatGenerationTask
|
|
from app.services.celery_download_recovery_service import (
|
|
ensure_aware_utc,
|
|
get_download_active_payloads,
|
|
get_due_download_record_ids,
|
|
postpone_download_active_check,
|
|
remove_download_active,
|
|
)
|
|
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
|
|
|
|
logger = logging.getLogger("video_gen")
|
|
|
|
|
|
def _now() -> datetime:
|
|
return datetime.now(timezone.utc)
|
|
|
|
|
|
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
|
|
checked = ensure_aware_utc(value)
|
|
if checked is None:
|
|
return True
|
|
return checked <= (now or _now())
|
|
|
|
|
|
def _queue_timeout_at(task: ChatGenerationTask, now: datetime | None = None) -> datetime:
|
|
current_time = now or _now()
|
|
enqueued_at = ensure_aware_utc(task.download_enqueued_at)
|
|
if enqueued_at is None:
|
|
return current_time
|
|
return enqueued_at + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
|
|
|
|
|
|
def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) -> bool:
|
|
current_time = now or _now()
|
|
return _queue_timeout_at(task, current_time) <= current_time
|
|
|
|
|
|
def _is_final_task_state(task: ChatGenerationTask) -> bool:
|
|
return task.status in ("completed", "failed") or task.pipeline_stage in (
|
|
"done",
|
|
"failed",
|
|
"timeout",
|
|
"download_failed",
|
|
)
|
|
|
|
|
|
async def recover_one_download_task(
|
|
db: AsyncSession,
|
|
task: ChatGenerationTask,
|
|
*,
|
|
payload: dict[str, Any] | None = None,
|
|
source: str = "startup_db",
|
|
) -> str:
|
|
from app.tasks.generation_download_tasks import (
|
|
DOWNLOAD_STAGE_DOWNLOADING,
|
|
DOWNLOAD_STAGE_QUEUED,
|
|
DOWNLOAD_STAGE_RETRY_WAITING,
|
|
enqueue_download_task,
|
|
)
|
|
|
|
current_time = _now()
|
|
|
|
if not task:
|
|
return "skip_missing_task"
|
|
if task.generation_mode != "chatapi_async":
|
|
await remove_download_active(task.id)
|
|
return "clean_invalid_mode"
|
|
if _is_final_task_state(task):
|
|
await remove_download_active(task.id)
|
|
return "clean_final_state"
|
|
if task.status != "generating":
|
|
await remove_download_active(task.id)
|
|
return "clean_not_generating"
|
|
if not task.remote_result_url:
|
|
return "skip_no_remote_result_url"
|
|
|
|
stage = task.pipeline_stage
|
|
redis_payload = payload or {}
|
|
|
|
if stage == "result_ready":
|
|
await log_task_event(
|
|
task,
|
|
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
|
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
|
|
detail={"payload": redis_payload},
|
|
)
|
|
await enqueue_download_task(
|
|
db,
|
|
task,
|
|
recover=True,
|
|
reason=f"{source}_result_ready",
|
|
)
|
|
return "recover_result_ready"
|
|
|
|
if stage == DOWNLOAD_STAGE_QUEUED:
|
|
if _is_queue_timeout(task, current_time):
|
|
await log_task_event(
|
|
task,
|
|
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
|
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
|
|
detail={"payload": redis_payload},
|
|
)
|
|
await enqueue_download_task(
|
|
db,
|
|
task,
|
|
recover=True,
|
|
reason=f"{source}_download_queued_timeout",
|
|
)
|
|
return "recover_queued_timeout"
|
|
|
|
await postpone_download_active_check(
|
|
record_id=task.id,
|
|
payload=payload,
|
|
check_at=_queue_timeout_at(task, current_time),
|
|
)
|
|
return "skip_queued_not_timeout"
|
|
|
|
if stage == DOWNLOAD_STAGE_DOWNLOADING:
|
|
if _is_expired(task.download_lease_until, current_time):
|
|
await log_task_event(
|
|
task,
|
|
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
|
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
|
|
detail={"payload": redis_payload},
|
|
)
|
|
await enqueue_download_task(
|
|
db,
|
|
task,
|
|
recover=True,
|
|
reason=f"{source}_downloading_lease_expired",
|
|
)
|
|
return "recover_downloading_expired"
|
|
|
|
await postpone_download_active_check(
|
|
record_id=task.id,
|
|
payload=payload,
|
|
check_at=task.download_lease_until,
|
|
)
|
|
return "skip_downloading_alive"
|
|
|
|
if stage == DOWNLOAD_STAGE_RETRY_WAITING:
|
|
if _is_expired(task.download_next_retry_at, current_time):
|
|
await log_task_event(
|
|
task,
|
|
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
|
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
|
|
detail={"payload": redis_payload},
|
|
)
|
|
await enqueue_download_task(
|
|
db,
|
|
task,
|
|
recover=True,
|
|
reason=f"{source}_retry_waiting_due",
|
|
)
|
|
return "recover_retry_due"
|
|
|
|
await postpone_download_active_check(
|
|
record_id=task.id,
|
|
payload=payload,
|
|
check_at=task.download_next_retry_at,
|
|
)
|
|
return "skip_retry_waiting_not_due"
|
|
|
|
return f"skip_stage_{stage}"
|
|
|
|
|
|
async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
|
"""启动时下载容灾扫描。
|
|
|
|
先按 Redis active_index 找到到期下载任务;Redis 不可用或索引丢失时,
|
|
再通过 DB fallback 扫描 result_ready/download_* 状态,避免任务永久卡住。
|
|
"""
|
|
checked_ids: set[str] = set()
|
|
results: dict[str, int] = {}
|
|
|
|
due_ids = await get_due_download_record_ids(
|
|
limit=settings.DOWNLOAD_RECOVERY_BATCH_SIZE,
|
|
)
|
|
payloads = await get_download_active_payloads(due_ids)
|
|
|
|
for task_id in due_ids:
|
|
result = await db.execute(
|
|
select(ChatGenerationTask)
|
|
.where(
|
|
ChatGenerationTask.id == task_id,
|
|
ChatGenerationTask.deleted_at.is_(None),
|
|
)
|
|
.with_for_update()
|
|
.limit(1)
|
|
)
|
|
task = result.scalar_one_or_none()
|
|
if task is None:
|
|
await remove_download_active(task_id)
|
|
action = "clean_missing_task"
|
|
else:
|
|
checked_ids.add(task.id)
|
|
action = await recover_one_download_task(
|
|
db,
|
|
task,
|
|
payload=payloads.get(task_id),
|
|
source="startup_redis",
|
|
)
|
|
results[action] = results.get(action, 0) + 1
|
|
|
|
# DB fallback:不依赖 Redis active 注册表。
|
|
fallback_result = await db.execute(
|
|
select(ChatGenerationTask)
|
|
.where(
|
|
ChatGenerationTask.deleted_at.is_(None),
|
|
ChatGenerationTask.generation_mode == "chatapi_async",
|
|
ChatGenerationTask.status == "generating",
|
|
ChatGenerationTask.remote_result_url.is_not(None),
|
|
ChatGenerationTask.pipeline_stage.in_(
|
|
["result_ready", "download_queued", "downloading", "retry_waiting"]
|
|
),
|
|
)
|
|
.order_by(ChatGenerationTask.updated_at.asc())
|
|
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
|
|
.with_for_update(skip_locked=True)
|
|
)
|
|
fallback_tasks = fallback_result.scalars().all()
|
|
|
|
for task in fallback_tasks:
|
|
if task.id in checked_ids:
|
|
continue
|
|
action = await recover_one_download_task(
|
|
db,
|
|
task,
|
|
payload=None,
|
|
source="startup_db",
|
|
)
|
|
results[action] = results.get(action, 0) + 1
|
|
checked_ids.add(task.id)
|
|
|
|
return {"checked": len(checked_ids), "results": results}
|
|
|
|
|
|
async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
|
"""启动时生成链路容灾扫描。
|
|
|
|
只在 Celery worker 启动时跑一次,不引入 beat,不新增第四条启动命令。
|
|
用于把 queued/creating/waiting_remote/polling/result_ready 等中间态重新投递到现有三个队列。
|
|
"""
|
|
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
|
from app.tasks.generation_download_tasks import enqueue_download_task
|
|
from app.tasks.generation_poll_tasks import poll_generation_task
|
|
|
|
current_time = _now()
|
|
results: dict[str, int] = {}
|
|
|
|
query_result = await db.execute(
|
|
select(ChatGenerationTask)
|
|
.where(
|
|
ChatGenerationTask.deleted_at.is_(None),
|
|
ChatGenerationTask.generation_mode == "chatapi_async",
|
|
ChatGenerationTask.status == "generating",
|
|
ChatGenerationTask.pipeline_stage.in_(
|
|
[
|
|
"queued",
|
|
"preparing",
|
|
"creating_provider_task",
|
|
"waiting_remote",
|
|
"polling",
|
|
"result_ready",
|
|
]
|
|
),
|
|
)
|
|
.order_by(ChatGenerationTask.updated_at.asc())
|
|
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
|
|
.with_for_update(skip_locked=True)
|
|
)
|
|
tasks = query_result.scalars().all()
|
|
|
|
for task in tasks:
|
|
if task.deadline_at and _is_expired(task.deadline_at, current_time):
|
|
await mark_chat_generation_task_failed_and_refund_once(
|
|
db,
|
|
task=task,
|
|
error_message="任务超时",
|
|
pipeline_stage="timeout",
|
|
)
|
|
await db.commit()
|
|
await log_task_event(
|
|
task,
|
|
event_type="TASK_TIMEOUT",
|
|
to_status="failed",
|
|
to_stage="timeout",
|
|
)
|
|
action = "mark_timeout"
|
|
|
|
elif task.pipeline_stage in ("queued", "preparing", "creating_provider_task"):
|
|
await log_task_event(
|
|
task,
|
|
event_type="GENERATION_RECOVERY_ENQUEUE",
|
|
message="启动时发现创建阶段任务未完成,恢复投递创建队列",
|
|
detail={"pipeline_stage": task.pipeline_stage},
|
|
)
|
|
chatapi_create_generation_task.apply_async(
|
|
args=[task.id],
|
|
queue="gen_chatapi_create",
|
|
countdown=0,
|
|
)
|
|
action = "recover_create"
|
|
|
|
elif task.pipeline_stage in ("waiting_remote", "polling"):
|
|
if task.remote_result_url:
|
|
await enqueue_download_task(
|
|
db,
|
|
task,
|
|
recover=True,
|
|
reason="startup_waiting_remote_has_result",
|
|
)
|
|
action = "recover_waiting_has_result"
|
|
elif task.provider_task_id or task.seedance_task_id:
|
|
await log_task_event(
|
|
task,
|
|
event_type="GENERATION_RECOVERY_ENQUEUE",
|
|
message="启动时发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
|
|
detail={"pipeline_stage": task.pipeline_stage},
|
|
)
|
|
task.pipeline_stage = "waiting_remote"
|
|
await db.commit()
|
|
poll_generation_task.apply_async(
|
|
args=[task.id],
|
|
queue="gen_provider_poll",
|
|
countdown=0,
|
|
)
|
|
action = "recover_poll"
|
|
else:
|
|
await log_task_event(
|
|
task,
|
|
event_type="GENERATION_RECOVERY_ENQUEUE",
|
|
message="启动时发现任务缺少供应商任务ID,恢复投递创建队列",
|
|
detail={"pipeline_stage": task.pipeline_stage},
|
|
)
|
|
task.pipeline_stage = "queued"
|
|
await db.commit()
|
|
chatapi_create_generation_task.apply_async(
|
|
args=[task.id],
|
|
queue="gen_chatapi_create",
|
|
countdown=0,
|
|
)
|
|
action = "recover_create_missing_provider_id"
|
|
|
|
elif task.pipeline_stage == "result_ready":
|
|
if task.remote_result_url:
|
|
await enqueue_download_task(
|
|
db,
|
|
task,
|
|
recover=True,
|
|
reason="startup_generation_result_ready",
|
|
)
|
|
action = "recover_result_ready"
|
|
else:
|
|
action = "skip_result_ready_no_url"
|
|
else:
|
|
action = f"skip_stage_{task.pipeline_stage}"
|
|
|
|
results[action] = results.get(action, 0) + 1
|
|
|
|
# 下载阶段单独跑 DB fallback。
|
|
download_result = await recover_download_tasks_once(db)
|
|
return {
|
|
"checked": len(tasks),
|
|
"results": results,
|
|
"download_recovery": download_result,
|
|
}
|