626 lines
24 KiB
Python
626 lines
24 KiB
Python
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}
|