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", 600) or 600)) 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}