Files
video-gen/video-gen-api/app/tasks/generation_recovery_tasks.py

546 lines
22 KiB
Python

# app/tasks/generation_recovery_tasks.py
from __future__ import annotations
import logging
from typing import Any, Awaitable, Callable, Dict
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.models.base import async_session
from app.services.celery_runtime.recovery_service import clear_startup_barrier, guard_periodic_recovery
from app.services.redis_registry_service import (
RedisExecutionLockLease,
get_registry_redis,
)
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or CeleryQueue.GEN_RECOVERY.value
RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
async def _recover_generation_records_once(*, include_create: bool, include_poll: bool, include_download: bool) -> Dict[str, Any]:
from app.services.generation.pipeline.recovery_repository import find_generation_record_recovery_batch
from app.tasks.generation_create_tasks import chatapi_create_generation_task
from app.tasks.generation_poll_tasks import poll_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
counts: dict[str, Any] = {
"create": 0,
"poll": 0,
"download": 0,
"inconsistent": 0,
"inconsistent_timeout": 0,
"errors": [],
}
batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20))
cursor = None
async with async_session() as db:
while True:
batch = await find_generation_record_recovery_batch(db, limit=batch_size, cursor=cursor)
if include_create or include_poll:
from datetime import datetime, timezone
from app.enums.generation_status import GenerationRecordPipelineStage
from app.enums.generation_task import ChatGenerationTaskEventType, GenerationMode
from app.services.generation.log_service import log_task_event
from app.services.generation.pipeline.owner_service import load_generation_owner
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.services.redis_registry_service import ensure_aware_utc
for ref in batch.inconsistent:
try:
owner = await load_generation_owner(
db,
owner_type=ref.owner_type,
owner_id=ref.owner_id,
for_update=True,
)
if (
owner is None
or int(owner.generation_attempt_no or 1)
!= int(ref.generation_attempt_no or 1)
):
await db.rollback()
continue
deadline_at = ensure_aware_utc(getattr(owner, "deadline_at", None))
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
previous_stage = str(owner.pipeline_stage or "")
if deadline_at is not None and deadline_at <= datetime.now(timezone.utc):
owner.pipeline_stage = GenerationRecordPipelineStage.TIMEOUT.value
await mark_generation_record_failed_and_refund_once(
db,
record=owner,
generation_attempt_no=ref.generation_attempt_no,
error_message="恢复证据异常且已超过任务截止时间",
)
await db.commit()
await log_task_event(
owner_type=ref.owner_type,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=GenerationMode.GENERATION_RECORD.value,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
from_stage=previous_stage,
to_stage=GenerationRecordPipelineStage.TIMEOUT.value,
message="恢复证据异常任务超过截止时间,已失败并幂等退款",
)
counts["inconsistent_timeout"] += 1
continue
if previous_stage != GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value:
owner.pipeline_stage = GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value
owner.error_message = (
f"恢复证据异常:阶段 {previous_stage} 缺少 remote_result_url 和供应商任务ID"
)
await db.commit()
await log_task_event(
owner_type=ref.owner_type,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=GenerationMode.GENERATION_RECORD.value,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_INCONSISTENT.value,
from_stage=previous_stage,
to_stage=GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
message="恢复证据异常,已隔离且不重新创建供应商任务",
)
else:
await db.rollback()
counts["inconsistent"] += 1
except Exception as exc:
await db.rollback()
counts["errors"].append(
{
"owner_id": ref.owner_id,
"stage": "recovery_inconsistent",
"error": str(exc),
}
)
if include_create:
for ref in batch.create:
try:
chatapi_create_generation_task.apply_async(
args=[ref.owner_id],
kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
task_id=f"generation-create:{ref.owner_type}:{ref.owner_id}:attempt:{ref.generation_attempt_no}",
)
counts["create"] += 1
except Exception as exc:
counts["errors"].append({"owner_id": ref.owner_id, "stage": "create", "error": str(exc)})
if include_poll:
for ref in batch.poll:
try:
poll_generation_task.apply_async(
args=[ref.owner_id],
kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no, "force_due": False},
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
)
counts["poll"] += 1
except Exception as exc:
counts["errors"].append({"owner_id": ref.owner_id, "stage": "poll", "error": str(exc)})
if include_download and batch.download:
from app.services.generation.pipeline.owner_service import load_generation_owner
for ref in batch.download:
try:
# 每条候选重新读取最新状态并只锁当前一行;前一条 commit 后
# 不继续使用批量查询得到的旧 ORM 对象。
owner = await load_generation_owner(
db,
owner_type=ref.owner_type,
owner_id=ref.owner_id,
for_update=True,
)
if (
owner is None
or int(owner.generation_attempt_no or 1)
!= int(ref.generation_attempt_no or 1)
):
await db.rollback()
continue
task_id = await enqueue_download_task(
db,
owner,
recover=True,
reason="generation_record_recovery",
)
if task_id:
counts["download"] += 1
except Exception as exc:
await db.rollback()
counts["errors"].append(
{
"owner_id": ref.owner_id,
"stage": "download",
"error": str(exc),
}
)
if batch.next_cursor is None:
break
cursor = batch.next_cursor
return counts
async def _run_download_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_download_tasks_once
async with async_session() as db:
chat_result = await recover_download_tasks_once(db)
record_result = await _recover_generation_records_once(include_create=False, include_poll=False, include_download=True)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_generation_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_generation_tasks_once
async with async_session() as db:
chat_result = await recover_generation_tasks_once(db)
record_result = await _recover_generation_records_once(include_create=True, include_poll=True, include_download=False)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_create_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_stale_create_tasks_once
async with async_session() as db:
chat_result = await recover_stale_create_tasks_once(db)
record_result = await _recover_generation_records_once(
include_create=True, include_poll=False, include_download=False
)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
from datetime import datetime, timedelta, timezone
from sqlalchemy import select
from app.enums.generation_status import GenerationRecordPipelineStage, GenerationStatus
from app.enums.generation_task import GenerationOwnerType
from app.models.generation_record import GenerationRecord
from app.services.generation.pipeline.db_lock_service import apply_short_lock_timeout
from app.services.generation.recovery_service import dispatch_due_poll_tasks_once
from app.tasks.generation_poll_tasks import poll_generation_task
async with async_session() as db:
chat_result = await dispatch_due_poll_tasks_once(db)
current_time = datetime.now(timezone.utc)
queue_hold_until = current_time + timedelta(seconds=int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120))
await apply_short_lock_timeout(db)
due_result = await db.execute(
select(GenerationRecord)
.where(
GenerationRecord.deleted_at.is_(None),
GenerationRecord.status == GenerationStatus.generating.value,
GenerationRecord.gen_type == "video",
GenerationRecord.pipeline_stage.in_([
GenerationRecordPipelineStage.WAITING_REMOTE.value,
GenerationRecordPipelineStage.POLLING.value,
]),
GenerationRecord.next_poll_at.is_not(None),
GenerationRecord.next_poll_at <= current_time,
)
.order_by(GenerationRecord.next_poll_at.asc(), GenerationRecord.id.asc())
.limit(int(settings.POLL_DUE_DISPATCH_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
owners = list(due_result.scalars().all())
dispatch_refs = []
for owner in owners:
dispatch_refs.append((str(owner.id), int(owner.generation_attempt_no or 1)))
owner.pipeline_stage = GenerationRecordPipelineStage.WAITING_REMOTE.value
owner.next_poll_at = queue_hold_until
await db.commit()
dispatched = 0
errors = []
for owner_id, attempt_no in dispatch_refs:
try:
poll_generation_task.apply_async(
args=[owner_id],
kwargs={
"owner_type": GenerationOwnerType.GENERATION_RECORD.value,
"generation_attempt_no": attempt_no,
"force_due": True,
},
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
)
dispatched += 1
except Exception as exc:
errors.append({"owner_id": owner_id, "error": str(exc)})
return {"chat_generation_task": chat_result, "generation_record": {"dispatched": dispatched, "errors": errors}}
async def _run_module_async_once() -> Dict[str, Any]:
from app.services.module_async_recovery_service import recover_module_async_tasks_once
async with async_session() as db:
return await recover_module_async_tasks_once(db)
async def _run_shot_split_once() -> Dict[str, Any]:
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
async with async_session() as db:
return await recover_shot_split_tasks_once(db)
async def _run_shot_analysis_once() -> Dict[str, Any]:
from app.services.shot_replicate_recovery_service import recover_shot_analysis_tasks_once
async with async_session() as db:
return await recover_shot_analysis_tasks_once(db)
async def _run_video_upscale_once() -> Dict[str, Any]:
from app.services.video_upscale.task_service import recover_video_upscale_tasks_once
async with async_session() as db:
return await recover_video_upscale_tasks_once(db)
async def _run_with_execution_lock(
*,
lock_key: str,
log_context: str,
runner: RecoveryRunner,
ttl_seconds: int | None = None,
) -> Dict[str, Any]:
"""恢复协调器严格依赖 Redis 执行锁,不允许 Redis 故障时无锁扫库。"""
lease = await RedisExecutionLockLease.acquire(
lock_key=lock_key,
ttl_seconds=int(ttl_seconds or settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context=log_context,
renew_interval_seconds=max(1, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30)),
)
if lease is None:
return {"skipped": "lock_held", "lock_key": lock_key}
try:
result = await runner()
await lease.ensure_owned()
result["execution_lock"] = "lock_acquired"
return result
finally:
await lease.close()
async def _is_lock_held(lock_key: str) -> bool:
redis = await get_registry_redis()
if redis is None:
return False
try:
return bool(await redis.exists(lock_key))
except Exception:
return False
async def _startup_or_generation_recovery_running() -> str | None:
# Beat 触发 dispatcher 时,如果启动容灾或完整生成容灾还在跑,直接跳过本轮。
# gen_recovery concurrency=1 已经能串行;这里是多机部署、残留消息、手动触发时的双保险。
lock_checks = [
("startup_recovery", settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY),
("generation_recovery", settings.GENERATION_RECOVERY_LOCK_KEY),
]
for name, lock_key in lock_checks:
if await _is_lock_held(lock_key):
return name
return None
async def _run_due_poll_dispatch_with_guard() -> Dict[str, Any]:
running = await _startup_or_generation_recovery_running()
if running:
return {"skipped": "recovery_lock_held", "lock": running}
return await _run_due_poll_dispatch_once()
async def _run_periodic_with_guard(
*,
lock_key: str,
log_context: str,
runner: RecoveryRunner,
ttl_seconds: int | None = None,
) -> Dict[str, Any]:
guarded = await guard_periodic_recovery()
if guarded:
return guarded
return await _run_with_execution_lock(
lock_key=lock_key,
log_context=log_context,
runner=runner,
ttl_seconds=ttl_seconds,
)
async def _run_startup_recovery_once() -> Dict[str, Any]:
"""启动容灾协调器;全局锁隔离周期对账,完成或失败后释放 barrier。"""
async def _run_startup_locked() -> Dict[str, Any]:
return await _run_with_execution_lock(
lock_key=settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY,
log_context="startup_recovery_once",
runner=_run_startup_recovery_steps,
)
try:
return await _run_with_execution_lock(
lock_key=settings.CELERY_RUNTIME_GLOBAL_RECOVERY_LOCK_KEY,
log_context="startup_global_recovery",
runner=_run_startup_locked,
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
)
finally:
await clear_startup_barrier()
async def _run_startup_recovery_steps() -> Dict[str, Any]:
results: Dict[str, Any] = {}
steps: list[tuple[str, str, str, RecoveryRunner]] = [
(
"module_async",
settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
"module_async_recovery",
_run_module_async_once,
),
(
"shot_analysis",
settings.SHOT_ANALYSIS_RECOVERY_LOCK_KEY,
"shot_analysis_recovery",
_run_shot_analysis_once,
),
(
"shot_split",
settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
"shot_split_recovery",
_run_shot_split_once,
),
(
"generation",
settings.GENERATION_RECOVERY_LOCK_KEY,
"generation_recovery",
_run_generation_once,
),
(
"download",
settings.DOWNLOAD_RECOVERY_LOCK_KEY,
"download_recovery",
_run_download_once,
),
(
"video_upscale",
settings.VIDEO_UPSCALE_RECOVERY_LOCK_KEY,
"video_upscale_recovery",
_run_video_upscale_once,
),
]
for name, lock_key, log_context, runner in steps:
try:
results[name] = await _run_with_execution_lock(
lock_key=lock_key,
log_context=log_context,
runner=runner,
)
except Exception as exc:
logger.exception("启动容灾步骤执行失败。step=%s", name)
results[name] = {"error": str(exc)}
return {"steps": results}
if celery_app:
@celery_app.task(
name="recovery.startup_recovery_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def startup_recovery_once(self) -> Dict[str, Any]:
return run_async(_run_startup_recovery_once())
@celery_app.task(
name="generation.recover_download_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_download_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_periodic_with_guard(
lock_key=settings.DOWNLOAD_RECOVERY_LOCK_KEY,
log_context="download_recovery",
runner=_run_download_once,
)
)
@celery_app.task(
name="generation.recover_generation_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_generation_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_periodic_with_guard(
lock_key=settings.GENERATION_RECOVERY_LOCK_KEY,
log_context="generation_recovery",
runner=_run_generation_once,
)
)
@celery_app.task(
name="generation.recover_create_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_create_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_periodic_with_guard(
lock_key=f"{settings.GENERATION_RECOVERY_LOCK_KEY}:create",
log_context="generation_create_recovery",
runner=_run_create_once,
ttl_seconds=max(55, int(settings.GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS or 60) - 5),
)
)
@celery_app.task(
name="generation.dispatch_due_poll_tasks",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def dispatch_due_poll_tasks(self) -> Dict[str, Any]:
return run_async(
_run_periodic_with_guard(
lock_key=settings.POLL_DUE_DISPATCH_LOCK_KEY,
log_context="due_poll_dispatch",
runner=_run_due_poll_dispatch_with_guard,
ttl_seconds=int(settings.POLL_DUE_DISPATCH_LOCK_TTL_SECONDS or 55),
)
)
else:
class _DisabledTask:
def delay(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
def apply_async(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask()
recover_create_tasks_once = _DisabledTask()
dispatch_due_poll_tasks = _DisabledTask()