Files
video-gen/video-gen-api/app/services/module_async_recovery_service.py

626 lines
24 KiB
Python
Raw Permalink 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 __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from typing import Any, Iterable
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.common import ModuleStepStatusEnum
from app.enums.hot_opening_replicate import HotOpeningStepCodeEnum, ModuleCodeEnum as HotModuleCodeEnum
from app.enums.shot_replicate import (
ModuleCodeEnum as ShotModuleCodeEnum,
ShotAnalysisStatusEnum,
ShotReplicateStepCodeEnum,
ShotSegmentAnalysisStatusEnum,
ShotSplitStatusEnum,
)
from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.shot_replicate_task_set import ShotReplicateTaskSet
from app.services.redis_registry_service import (
datetime_to_epoch,
redis_acquire_lock,
redis_get_due_registry_ids,
redis_get_registry_payloads,
redis_postpone_registry_item,
redis_release_lock,
redis_remove_registry_item,
redis_upsert_registry_item,
utc_now,
)
from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
QUEUE_CREATE = "gen_chatapi_create"
QUEUE_DOWNLOAD = "gen_result_download"
OBJECT_MODULE_STEP = "module_step"
OBJECT_SHOT_TASK_SET_ANALYSIS = "shot_task_set_analysis"
OBJECT_SHOT_SEGMENT_ANALYSIS = "shot_segment_analysis"
OBJECT_SHOT_SPLIT_SEGMENT = "shot_split_segment"
TASK_HOT_IMAGE_PROMPT = "hot_opening.start_image_prompt_optimize"
TASK_HOT_VIDEO_PROMPT = "hot_opening.start_video_prompt_optimize"
TASK_SHOT_IMAGE_PROMPT = "shot_replicate.start_image_prompt_optimize"
TASK_SHOT_VIDEO_PROMPT = "shot_replicate.start_video_prompt_optimize"
TASK_SHOT_ANALYZE_ORIGINAL = "shot_replicate.analyze_original_video"
TASK_SHOT_ANALYZE_CUSTOM_SEGMENT = "shot_replicate.analyze_custom_segment_video"
TASK_SHOT_SPLIT_ONE = "shot_replicate.split_one_segment"
HOT_MODULE = HotModuleCodeEnum.HOT_OPENING_REPLICATE.value
SHOT_MODULE = ShotModuleCodeEnum.SHOT_REPLICATE.value
TERMINAL_STEP_STATUSES = {
ModuleStepStatusEnum.COMPLETED.value,
ModuleStepStatusEnum.FAILED.value,
ModuleStepStatusEnum.CANCELLED.value,
}
TERMINAL_ANALYSIS_STATUSES = {
ShotAnalysisStatusEnum.COMPLETED.value,
ShotAnalysisStatusEnum.FAILED.value,
}
TERMINAL_SEGMENT_ANALYSIS_STATUSES = {
ShotSegmentAnalysisStatusEnum.COMPLETED.value,
ShotSegmentAnalysisStatusEnum.FAILED.value,
ShotSegmentAnalysisStatusEnum.NOT_REQUIRED.value,
}
TERMINAL_SPLIT_STATUSES = {
ShotSplitStatusEnum.COMPLETED.value,
ShotSplitStatusEnum.FAILED.value,
}
def _now() -> datetime:
return datetime.now(timezone.utc)
def _ensure_aware(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def _hash_key() -> str:
return settings.MODULE_ASYNC_ACTIVE_REDIS_HASH_KEY
def _zset_key() -> str:
return settings.MODULE_ASYNC_ACTIVE_REDIS_ZSET_KEY
def _lease_seconds() -> int:
return max(1, int(settings.MODULE_ASYNC_LEASE_SECONDS or 600))
def _shot_analysis_lease_seconds() -> int:
timeout = max(1, int(getattr(settings, "SHOT_ANALYSIS_TIMEOUT_SECONDS", 3600) or 3600))
return max(_lease_seconds(), timeout + 120)
def _queue_timeout_seconds() -> int:
return max(1, int(settings.MODULE_ASYNC_QUEUE_TIMEOUT_SECONDS or 300))
def _requeue_delay_seconds() -> int:
return max(1, int(settings.MODULE_ASYNC_REQUEUE_DELAY_SECONDS or 10))
def _item_id(object_type: str, object_id: str) -> str:
return f"{object_type}:{object_id}"
def _lock_key(object_type: str, object_id: str) -> str:
return f"{settings.MODULE_ASYNC_LOCK_KEY_PREFIX}:{object_type}:{object_id}"
def _check_at_after(seconds: int | None = None) -> datetime:
return _now() + timedelta(seconds=int(seconds or _lease_seconds()))
def _base_payload(
*,
object_type: str,
object_id: str,
task_name: str,
queue: str,
args: list[Any],
module: str | None = None,
project_id: str | None = None,
step_id: str | None = None,
step_code: str | None = None,
task_set_id: str | None = None,
segment_id: str | None = None,
reason: str = "submit",
) -> dict[str, Any]:
now_epoch = datetime_to_epoch(utc_now())
return {
"object_type": object_type,
"object_id": object_id,
"task_name": task_name,
"queue": queue,
"args": args,
"module": module,
"project_id": project_id,
"step_id": step_id,
"step_code": step_code,
"task_set_id": task_set_id,
"segment_id": segment_id,
"reason": reason,
"created_at": now_epoch,
"updated_at": now_epoch,
}
async def register_active_task(
*,
object_type: str,
object_id: str,
task_name: str,
queue: str,
args: list[Any],
module: str | None = None,
project_id: str | None = None,
step_id: str | None = None,
step_code: str | None = None,
task_set_id: str | None = None,
segment_id: str | None = None,
check_after_seconds: int | None = None,
reason: str = "submit",
) -> str:
item_id = _item_id(object_type, object_id)
payload = _base_payload(
object_type=object_type,
object_id=object_id,
task_name=task_name,
queue=queue,
args=args,
module=module,
project_id=project_id,
step_id=step_id,
step_code=step_code,
task_set_id=task_set_id,
segment_id=segment_id,
reason=reason,
)
await redis_upsert_registry_item(
hash_key=_hash_key(),
zset_key=_zset_key(),
item_id=item_id,
payload=payload,
check_at=_check_at_after(check_after_seconds),
log_context="module_async_active",
)
return item_id
async def register_module_step_task(
*,
module: str,
project_id: str,
step_id: str,
step_code: str,
task_name: str,
queue: str = QUEUE_CREATE,
) -> str:
return await register_active_task(
object_type=OBJECT_MODULE_STEP,
object_id=step_id,
task_name=task_name,
queue=queue,
args=[project_id, step_id],
module=module,
project_id=project_id,
step_id=step_id,
step_code=step_code,
check_after_seconds=_lease_seconds(),
)
async def register_shot_task_set_analysis_task(task_set_id: str) -> str:
return await register_active_task(
object_type=OBJECT_SHOT_TASK_SET_ANALYSIS,
object_id=task_set_id,
task_name=TASK_SHOT_ANALYZE_ORIGINAL,
queue=QUEUE_CREATE,
args=[task_set_id],
module=SHOT_MODULE,
project_id=task_set_id,
task_set_id=task_set_id,
check_after_seconds=_shot_analysis_lease_seconds(),
)
async def register_shot_segment_analysis_task(segment_id: str, *, task_set_id: str | None = None) -> str:
return await register_active_task(
object_type=OBJECT_SHOT_SEGMENT_ANALYSIS,
object_id=segment_id,
task_name=TASK_SHOT_ANALYZE_CUSTOM_SEGMENT,
queue=QUEUE_CREATE,
args=[segment_id],
module=SHOT_MODULE,
project_id=task_set_id,
task_set_id=task_set_id,
segment_id=segment_id,
check_after_seconds=_shot_analysis_lease_seconds(),
)
async def register_shot_split_task(segment_id: str, *, task_set_id: str | None = None) -> str:
return await register_active_task(
object_type=OBJECT_SHOT_SPLIT_SEGMENT,
object_id=segment_id,
task_name=TASK_SHOT_SPLIT_ONE,
queue=QUEUE_DOWNLOAD,
args=[segment_id],
module=SHOT_MODULE,
project_id=task_set_id,
task_set_id=task_set_id,
segment_id=segment_id,
check_after_seconds=max(_lease_seconds(), int(settings.SHOT_SPLIT_LEASE_SECONDS or 600)),
)
async def remove_active_task(*, object_type: str, object_id: str) -> None:
await redis_remove_registry_item(
hash_key=_hash_key(),
zset_key=_zset_key(),
item_id=_item_id(object_type, object_id),
log_context="module_async_active",
)
async def postpone_active_task(
*,
object_type: str,
object_id: str,
delay_seconds: int | None = None,
reason: str | None = None,
) -> None:
item_id = _item_id(object_type, object_id)
payloads = await redis_get_registry_payloads(
hash_key=_hash_key(),
item_ids=[item_id],
log_context="module_async_active",
)
payload = payloads.get(item_id)
if payload is not None and reason:
payload["reason"] = reason
await redis_postpone_registry_item(
hash_key=_hash_key(),
zset_key=_zset_key(),
item_id=item_id,
payload=payload,
check_at=_check_at_after(delay_seconds or _requeue_delay_seconds()),
log_context="module_async_active",
)
async def mark_active_started(*, object_type: str, object_id: str, reason: str = "started") -> None:
delay_seconds = _lease_seconds()
if object_type in {OBJECT_SHOT_TASK_SET_ANALYSIS, OBJECT_SHOT_SEGMENT_ANALYSIS}:
delay_seconds = _shot_analysis_lease_seconds()
elif object_type == OBJECT_SHOT_SPLIT_SEGMENT:
delay_seconds = max(_lease_seconds(), int(settings.SHOT_SPLIT_LEASE_SECONDS or 600))
await postpone_active_task(
object_type=object_type,
object_id=object_id,
delay_seconds=delay_seconds,
reason=reason,
)
async def acquire_object_lock(*, object_type: str, object_id: str) -> str | None:
return await redis_acquire_lock(
lock_key=_lock_key(object_type, object_id),
ttl_seconds=int(settings.MODULE_ASYNC_LOCK_TTL_SECONDS or _lease_seconds()),
log_context="module_async_object_lock",
)
async def release_object_lock(*, object_type: str, object_id: str, token: str | None) -> None:
if token:
await redis_release_lock(
lock_key=_lock_key(object_type, object_id),
token=token,
log_context="module_async_object_lock",
)
async def cleanup_active_if_terminal(db: AsyncSession, *, object_type: str, object_id: str) -> bool:
if object_type == OBJECT_MODULE_STEP:
result = await db.execute(
select(ModuleGenerationStep)
.where(ModuleGenerationStep.id == object_id)
.limit(1)
)
step = result.scalar_one_or_none()
if not step or step.deleted_at is not None or step.status in TERMINAL_STEP_STATUSES:
await remove_active_task(object_type=object_type, object_id=object_id)
return True
return False
if object_type == OBJECT_SHOT_TASK_SET_ANALYSIS:
result = await db.execute(
select(ShotReplicateTaskSet)
.where(ShotReplicateTaskSet.id == object_id)
.limit(1)
)
task_set = result.scalar_one_or_none()
if not task_set or task_set.deleted_at is not None or task_set.analysis_status in TERMINAL_ANALYSIS_STATUSES:
await remove_active_task(object_type=object_type, object_id=object_id)
return True
return False
if object_type == OBJECT_SHOT_SEGMENT_ANALYSIS:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == object_id)
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment or segment.deleted_at is not None or segment.analysis_status in TERMINAL_SEGMENT_ANALYSIS_STATUSES:
await remove_active_task(object_type=object_type, object_id=object_id)
return True
return False
if object_type == OBJECT_SHOT_SPLIT_SEGMENT:
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == object_id)
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment or segment.deleted_at is not None or segment.split_status in TERMINAL_SPLIT_STATUSES:
await remove_active_task(object_type=object_type, object_id=object_id)
return True
return False
return False
def _payload_args(payload: dict[str, Any]) -> list[Any]:
args = payload.get("args")
if isinstance(args, list):
return args
object_type = str(payload.get("object_type") or "")
if object_type == OBJECT_MODULE_STEP:
return [payload.get("project_id"), payload.get("step_id")]
if object_type == OBJECT_SHOT_TASK_SET_ANALYSIS:
return [payload.get("task_set_id") or payload.get("object_id")]
if object_type in {OBJECT_SHOT_SEGMENT_ANALYSIS, OBJECT_SHOT_SPLIT_SEGMENT}:
return [payload.get("segment_id") or payload.get("object_id")]
return []
def _send_task(task_name: str, *, args: list[Any], queue: str, countdown: int = 0, priority: int | None = None) -> bool:
if celery_app is None:
return False
if not task_name or not queue:
return False
celery_app.send_task(
task_name,
args=args,
queue=queue,
countdown=max(0, int(countdown or 0)),
priority=priority if priority is not None else settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
return True
async def _recover_payload_from_redis(db: AsyncSession, item_id: str, payload: dict[str, Any]) -> str:
object_type = str(payload.get("object_type") or "")
object_id = str(payload.get("object_id") or "")
task_name = str(payload.get("task_name") or "")
queue = str(payload.get("queue") or QUEUE_CREATE)
args = _payload_args(payload)
if not object_type or not object_id or not task_name or not args:
await redis_remove_registry_item(hash_key=_hash_key(), zset_key=_zset_key(), item_id=item_id, log_context="module_async_active")
return "remove_invalid_payload"
if await cleanup_active_if_terminal(db, object_type=object_type, object_id=object_id):
return "remove_terminal"
if object_type == OBJECT_SHOT_SPLIT_SEGMENT:
from app.services.shot_replicate_recovery_service import recover_one_split_segment
result = await db.execute(
select(ShotReplicateSegment)
.where(ShotReplicateSegment.id == object_id, ShotReplicateSegment.deleted_at.is_(None))
.with_for_update(skip_locked=True)
.limit(1)
)
segment = result.scalar_one_or_none()
if not segment:
await remove_active_task(object_type=object_type, object_id=object_id)
return "remove_missing_split_segment"
action = await recover_one_split_segment(db, segment, source="redis_active")
if action.startswith("recover_"):
await postpone_active_task(object_type=object_type, object_id=object_id, delay_seconds=int(settings.SHOT_SPLIT_LEASE_SECONDS or _lease_seconds()), reason="redis_recovered")
elif action.startswith("skip_completed") or action.startswith("skip_failed") or action.startswith("mark_failed"):
await remove_active_task(object_type=object_type, object_id=object_id)
return f"split_{action}"
try:
_send_task(task_name, args=args, queue=queue, countdown=0, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER)
await postpone_active_task(object_type=object_type, object_id=object_id, delay_seconds=_lease_seconds(), reason="redis_requeued")
return "redis_requeued"
except Exception as exc:
logger.exception("Redis active 恢复投递失败。item_id=%s, task=%s", item_id, task_name)
await postpone_active_task(object_type=object_type, object_id=object_id, delay_seconds=_requeue_delay_seconds(), reason=f"redis_requeue_failed:{exc}")
return "redis_requeue_failed"
async def _recover_due_redis_items(db: AsyncSession, *, limit: int) -> dict[str, int]:
item_ids = await redis_get_due_registry_ids(
zset_key=_zset_key(),
limit=limit,
now=_now(),
log_context="module_async_active",
)
if not item_ids:
return {}
payloads = await redis_get_registry_payloads(hash_key=_hash_key(), item_ids=item_ids, log_context="module_async_active")
results: dict[str, int] = {}
for item_id in item_ids:
payload = payloads.get(item_id)
if not payload:
await redis_remove_registry_item(hash_key=_hash_key(), zset_key=_zset_key(), item_id=item_id, log_context="module_async_active")
action = "remove_missing_payload"
else:
action = await _recover_payload_from_redis(db, item_id, payload)
results[action] = results.get(action, 0) + 1
return results
def _is_stale_datetime(value: datetime | None, *, seconds: int, now: datetime) -> bool:
checked = _ensure_aware(value)
if checked is None:
return True
return checked + timedelta(seconds=max(1, int(seconds))) <= now
async def _recover_stale_module_steps(db: AsyncSession, *, limit: int) -> dict[str, int]:
now = _now()
stale_cutoff = now - timedelta(seconds=_lease_seconds())
result = await db.execute(
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.status == ModuleStepStatusEnum.PROCESSING.value,
ModuleGenerationStep.module.in_([HOT_MODULE, SHOT_MODULE]),
ModuleGenerationStep.step_code.in_(
[
HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
]
),
ModuleGenerationStep.updated_at <= stale_cutoff,
)
.order_by(ModuleGenerationStep.updated_at.asc())
.limit(limit)
.with_for_update(skip_locked=True)
)
steps = list(result.scalars().all())
results: dict[str, int] = {}
for step in steps:
if step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value:
task_name = TASK_HOT_IMAGE_PROMPT
elif step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
task_name = TASK_HOT_VIDEO_PROMPT
elif step.module == SHOT_MODULE and step.step_code == ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value:
task_name = TASK_SHOT_IMAGE_PROMPT
elif step.module == SHOT_MODULE and step.step_code == ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
task_name = TASK_SHOT_VIDEO_PROMPT
else:
results["skip_unknown_step"] = results.get("skip_unknown_step", 0) + 1
continue
await register_module_step_task(
module=step.module,
project_id=step.project_id,
step_id=step.id,
step_code=step.step_code,
task_name=task_name,
queue=QUEUE_CREATE,
)
try:
_send_task(task_name, args=[step.project_id, step.id], queue=QUEUE_CREATE, countdown=0, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER)
results["db_step_requeued"] = results.get("db_step_requeued", 0) + 1
except Exception:
logger.exception("DB fallback 恢复模块步骤失败。step_id=%s", step.id)
results["db_step_requeue_failed"] = results.get("db_step_requeue_failed", 0) + 1
await db.commit()
return results
async def _recover_stale_shot_task_sets(db: AsyncSession, *, limit: int) -> dict[str, int]:
now = _now()
stale_cutoff = now - timedelta(seconds=_lease_seconds())
result = await db.execute(
select(ShotReplicateTaskSet)
.where(
ShotReplicateTaskSet.deleted_at.is_(None),
ShotReplicateTaskSet.analysis_status.in_([ShotAnalysisStatusEnum.PENDING.value, ShotAnalysisStatusEnum.PROCESSING.value]),
ShotReplicateTaskSet.updated_at <= stale_cutoff,
)
.order_by(ShotReplicateTaskSet.updated_at.asc())
.limit(limit)
.with_for_update(skip_locked=True)
)
task_sets = list(result.scalars().all())
results: dict[str, int] = {}
for task_set in task_sets:
await register_shot_task_set_analysis_task(task_set.id)
try:
_send_task(TASK_SHOT_ANALYZE_ORIGINAL, args=[task_set.id], queue=QUEUE_CREATE, countdown=0, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER)
results["db_task_set_analysis_requeued"] = results.get("db_task_set_analysis_requeued", 0) + 1
except Exception:
logger.exception("DB fallback 恢复拆镜原视频分析失败。task_set_id=%s", task_set.id)
results["db_task_set_analysis_requeue_failed"] = results.get("db_task_set_analysis_requeue_failed", 0) + 1
await db.commit()
return results
async def _recover_stale_shot_segment_analysis(db: AsyncSession, *, limit: int) -> dict[str, int]:
now = _now()
stale_cutoff = now - timedelta(seconds=_lease_seconds())
result = await db.execute(
select(ShotReplicateSegment)
.where(
ShotReplicateSegment.deleted_at.is_(None),
ShotReplicateSegment.segment_video_url.is_not(None),
ShotReplicateSegment.analysis_status.in_([ShotSegmentAnalysisStatusEnum.PENDING.value, ShotSegmentAnalysisStatusEnum.PROCESSING.value]),
ShotReplicateSegment.updated_at <= stale_cutoff,
)
.order_by(ShotReplicateSegment.updated_at.asc())
.limit(limit)
.with_for_update(skip_locked=True)
)
segments = list(result.scalars().all())
results: dict[str, int] = {}
for segment in segments:
await register_shot_segment_analysis_task(segment.id, task_set_id=segment.task_set_id)
try:
_send_task(TASK_SHOT_ANALYZE_CUSTOM_SEGMENT, args=[segment.id], queue=QUEUE_CREATE, countdown=0, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER)
results["db_segment_analysis_requeued"] = results.get("db_segment_analysis_requeued", 0) + 1
except Exception:
logger.exception("DB fallback 恢复拆镜片段分析失败。segment_id=%s", segment.id)
results["db_segment_analysis_requeue_failed"] = results.get("db_segment_analysis_requeue_failed", 0) + 1
await db.commit()
return results
def _merge_counts(target: dict[str, int], items: Iterable[tuple[str, int]]) -> None:
for key, value in items:
target[key] = target.get(key, 0) + int(value)
async def recover_module_async_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""统一恢复模块异步任务。\n\n 覆盖范围:\n - hot_opening / shot_replicate 的图片、视频 AI 提词步骤;\n - shot_replicate 原视频分析;\n - shot_replicate 自定义片段分析;\n - shot_replicate split active 注册项。\n\n 拆镜 split 的 DB fallback 仍保留在 shot_replicate_recovery_service\n 这里主要补 Redis active 恢复和非 split 类任务的 DB fallback。\n """
batch_size = max(1, int(settings.MODULE_ASYNC_RECOVERY_BATCH_SIZE or 100))
results: dict[str, int] = {}
redis_results = await _recover_due_redis_items(db, limit=batch_size)
_merge_counts(results, redis_results.items())
step_results = await _recover_stale_module_steps(db, limit=batch_size)
_merge_counts(results, step_results.items())
task_set_results = await _recover_stale_shot_task_sets(db, limit=batch_size)
_merge_counts(results, task_set_results.items())
segment_results = await _recover_stale_shot_segment_analysis(db, limit=batch_size)
_merge_counts(results, segment_results.items())
return {"checked": sum(results.values()), "results": results}