diff --git a/video-gen-api/app/api/v1/hot_opening_replicate.py b/video-gen-api/app/api/v1/hot_opening_replicate.py index 5704d42d..c7c9b493 100644 --- a/video-gen-api/app/api/v1/hot_opening_replicate.py +++ b/video-gen-api/app/api/v1/hot_opening_replicate.py @@ -8,7 +8,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.dependencies import get_current_user, get_db from app.models.user import User -from app.enums.hot_opening_replicate import ModuleCodeEnum +from app.enums.hot_opening_replicate import HotOpeningStepCodeEnum, ModuleCodeEnum from app.schemas.hot_opening_replicate import ( HotOpeningActionOut, HotOpeningDeleteOut, @@ -40,6 +40,11 @@ from app.services.hot_opening_replicate_service import ( update_hot_opening_video_prompt_schema, ) from app.services.module_generation_log_service import log_module_error +from app.services.module_async_recovery_service import ( + TASK_HOT_IMAGE_PROMPT, + TASK_HOT_VIDEO_PROMPT, + register_module_step_task, +) from app.tasks.celery_app import celery_app MODULE = ModuleCodeEnum.HOT_OPENING_REPLICATE.value @@ -431,8 +436,15 @@ async def generate_image_prompt( from app.tasks.hot_opening_replicate_tasks import start_image_prompt_optimize + await register_module_step_task( + module=MODULE, + project_id=project_id_value, + step_id=step_id_value, + step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + task_name=TASK_HOT_IMAGE_PROMPT, + ) try: - start_image_prompt_optimize.delay(project_id_value, step_id_value) + start_image_prompt_optimize.apply_async(args=[project_id_value, step_id_value], queue="gen_chatapi_create", countdown=0) except Exception as exc: await _mark_dispatch_failed_and_raise( db, @@ -560,8 +572,15 @@ async def generate_video_prompt( from app.tasks.hot_opening_replicate_tasks import start_video_prompt_optimize + await register_module_step_task( + module=MODULE, + project_id=project_id_value, + step_id=step_id_value, + step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + task_name=TASK_HOT_VIDEO_PROMPT, + ) try: - start_video_prompt_optimize.delay(project_id_value, step_id_value) + start_video_prompt_optimize.apply_async(args=[project_id_value, step_id_value], queue="gen_chatapi_create", countdown=0) except Exception as exc: await _mark_dispatch_failed_and_raise( db, diff --git a/video-gen-api/app/api/v1/shot_replicate.py b/video-gen-api/app/api/v1/shot_replicate.py index f54095ed..d620f104 100644 --- a/video-gen-api/app/api/v1/shot_replicate.py +++ b/video-gen-api/app/api/v1/shot_replicate.py @@ -8,7 +8,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.dependencies import get_current_user, get_db from app.models.user import User -from app.enums.shot_replicate import ModuleCodeEnum +from app.enums.shot_replicate import ModuleCodeEnum, ShotReplicateStepCodeEnum from app.schemas.shot_replicate import ( ShotReplicateActionOut, ShotReplicateDeleteOut, @@ -57,6 +57,14 @@ from app.services.shot_replicate_taskset_service import ( task_set_detail, ) from app.services.module_generation_log_service import log_module_error, log_module_event_file +from app.services.module_async_recovery_service import ( + TASK_SHOT_IMAGE_PROMPT, + TASK_SHOT_VIDEO_PROMPT, + register_module_step_task, + register_shot_segment_analysis_task, + register_shot_split_task, + register_shot_task_set_analysis_task, +) from app.tasks.celery_app import celery_app MODULE = ModuleCodeEnum.SHOT_REPLICATE.value @@ -141,6 +149,22 @@ def _log_api_exception_from_locals(exc: BaseException, local_values: dict, messa exc=exc, ) + + +def _ensure_celery_enabled(*, current_user: User | None = None, project_id: str | None = None, step_id: str | None = None) -> None: + if celery_app is not None: + return + message = "Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker" + _log_api_error( + event_type="CELERY_DISABLED", + current_user=current_user, + project_id=project_id, + step_id=step_id, + message=message, + detail={"api": "shot_replicate"}, + ) + raise HTTPException(status_code=503, detail=message) + async def _reload_project_detail(db: AsyncSession, current_user: User, project_id: str) -> ShotReplicateTaskDetailOut: project = await _get_project_for_user( db, @@ -213,6 +237,7 @@ async def create_shot_task_set( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: task_set = await create_task_set(db, current_user=current_user, req=req) task_set_id = task_set.id @@ -225,22 +250,21 @@ async def create_shot_task_set( _log_api_exception_from_locals(exc, locals(), f"创建拆镜总任务集失败: {exc}") raise HTTPException(status_code=500, detail=f"创建拆镜总任务集失败: {exc}") - if celery_app: - try: - from app.tasks.shot_replicate_tasks import analyze_original_video + try: + from app.tasks.shot_replicate_tasks import analyze_original_video - analyze_original_video.apply_async(args=[task_set_id], queue="gen_chatapi_create", countdown=0) - except Exception as exc: - # 分析任务投递失败时保留总任务,前端可稍后通过恢复/重试处理。 - _log_api_error( - event_type="CELERY_DISPATCH_FAILED", - current_user=current_user, - project_id=task_set_id, - message=f"拆镜分析任务投递失败: {exc}", - detail={"task_set_id": task_set_id, "task": "analyze_original_video"}, - exc=exc, - ) - raise HTTPException(status_code=503, detail=f"拆镜分析任务投递失败: {exc}") + await register_shot_task_set_analysis_task(task_set_id) + analyze_original_video.apply_async(args=[task_set_id], queue="gen_chatapi_create", countdown=0) + except Exception as exc: + _log_api_error( + event_type="CELERY_DISPATCH_FAILED", + current_user=current_user, + project_id=task_set_id, + message=f"拆镜分析任务投递失败: {exc}", + detail={"task_set_id": task_set_id, "task": "analyze_original_video"}, + exc=exc, + ) + raise HTTPException(status_code=503, detail=f"拆镜分析任务投递失败: {exc}") return await task_set_detail(db, current_user=_user_context(current_user), task_set_id=task_set_id) @@ -296,6 +320,7 @@ async def split_by_ai( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: out = await create_segments_by_ai(db, current_user=current_user, task_set_id=task_set_id, req=req) segment_ids = [item.id for item in out.segments] @@ -308,11 +333,11 @@ async def split_by_ai( _log_api_exception_from_locals(exc, locals(), f"按 AI 建议拆镜失败: {exc}") raise HTTPException(status_code=500, detail=f"按 AI 建议拆镜失败: {exc}") - if celery_app: - from app.tasks.shot_replicate_tasks import split_one_segment + from app.tasks.shot_replicate_tasks import split_one_segment - for segment_id in segment_ids: - split_one_segment.apply_async(args=[segment_id], queue="gen_result_download", countdown=0) + for segment_id in segment_ids: + await register_shot_split_task(segment_id, task_set_id=task_set_id) + split_one_segment.apply_async(args=[segment_id], queue="gen_result_download", countdown=0) return out @@ -327,6 +352,7 @@ async def split_custom( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: out = await create_custom_segment(db, current_user=current_user, task_set_id=task_set_id, req=req) segment_id = out.segment.id @@ -339,10 +365,10 @@ async def split_custom( _log_api_exception_from_locals(exc, locals(), f"自定义拆镜失败: {exc}") raise HTTPException(status_code=500, detail=f"自定义拆镜失败: {exc}") - if celery_app: - from app.tasks.shot_replicate_tasks import split_one_segment + from app.tasks.shot_replicate_tasks import split_one_segment - split_one_segment.apply_async(args=[segment_id], queue="gen_result_download", countdown=0) + await register_shot_split_task(segment_id, task_set_id=task_set_id) + split_one_segment.apply_async(args=[segment_id], queue="gen_result_download", countdown=0) return out @@ -520,6 +546,7 @@ async def generate_image_prompt( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: project, step = await submit_image_prompt_optimize(db, current_user=current_user, project_id=project_id, req=req) project_id_value, step_id_value = project.id, step.id @@ -533,10 +560,16 @@ async def generate_image_prompt( raise HTTPException(status_code=500, detail=f"提交图片 AI 提词失败: {exc}") try: - if celery_app: - from app.tasks.shot_replicate_flow_tasks import start_image_prompt_optimize + from app.tasks.shot_replicate_flow_tasks import start_image_prompt_optimize - start_image_prompt_optimize.apply_async(args=[project_id_value, step_id_value], queue="gen_chatapi_create", countdown=0) + await register_module_step_task( + module=MODULE, + project_id=project_id_value, + step_id=step_id_value, + step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + task_name=TASK_SHOT_IMAGE_PROMPT, + ) + start_image_prompt_optimize.apply_async(args=[project_id_value, step_id_value], queue="gen_chatapi_create", countdown=0) except Exception as exc: await _mark_dispatch_failed_and_raise(db, current_user=current_user, project_id=project_id_value, step_id=step_id_value, message=f"图片 AI 提词任务投递失败: {exc}") @@ -554,6 +587,7 @@ async def generate_image( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: project, step, chat_task = await generate_image_from_prompt(db, current_user=current_user, project_id=project_id, req=req) project_id_value, step_id_value, chat_task_id_value = project.id, step.id, chat_task.id @@ -567,10 +601,9 @@ async def generate_image( raise HTTPException(status_code=500, detail=f"提交图片生成失败: {exc}") try: - if celery_app: - from app.tasks.generation_create_tasks import chatapi_create_generation_task + from app.tasks.generation_create_tasks import chatapi_create_generation_task - chatapi_create_generation_task.apply_async(args=[chat_task_id_value], queue="gen_chatapi_create", countdown=0) + chatapi_create_generation_task.apply_async(args=[chat_task_id_value], queue="gen_chatapi_create", countdown=0) except Exception as exc: await _mark_dispatch_failed_and_raise(db, current_user=current_user, project_id=project_id_value, step_id=step_id_value, message=f"图片生成任务投递失败: {exc}") @@ -588,6 +621,7 @@ async def generate_video_prompt( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: project, step = await submit_video_prompt_optimize(db, current_user=current_user, project_id=project_id, req=req) project_id_value, step_id_value = project.id, step.id @@ -601,10 +635,16 @@ async def generate_video_prompt( raise HTTPException(status_code=500, detail=f"提交视频 AI 提词失败: {exc}") try: - if celery_app: - from app.tasks.shot_replicate_flow_tasks import start_video_prompt_optimize + from app.tasks.shot_replicate_flow_tasks import start_video_prompt_optimize - start_video_prompt_optimize.apply_async(args=[project_id_value, step_id_value], queue="gen_chatapi_create", countdown=0) + await register_module_step_task( + module=MODULE, + project_id=project_id_value, + step_id=step_id_value, + step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + task_name=TASK_SHOT_VIDEO_PROMPT, + ) + start_video_prompt_optimize.apply_async(args=[project_id_value, step_id_value], queue="gen_chatapi_create", countdown=0) except Exception as exc: await _mark_dispatch_failed_and_raise(db, current_user=current_user, project_id=project_id_value, step_id=step_id_value, message=f"视频 AI 提词任务投递失败: {exc}") @@ -622,6 +662,7 @@ async def generate_video( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + _ensure_celery_enabled(current_user=current_user, project_id=locals().get("project_id") or locals().get("task_set_id")) try: project, step, chat_task = await generate_video_from_prompt(db, current_user=current_user, project_id=project_id, req=req) project_id_value, step_id_value, chat_task_id_value = project.id, step.id, chat_task.id @@ -635,10 +676,9 @@ async def generate_video( raise HTTPException(status_code=500, detail=f"提交视频生成失败: {exc}") try: - if celery_app: - from app.tasks.generation_create_tasks import chatapi_create_generation_task + from app.tasks.generation_create_tasks import chatapi_create_generation_task - chatapi_create_generation_task.apply_async(args=[chat_task_id_value], queue="gen_chatapi_create", countdown=0) + chatapi_create_generation_task.apply_async(args=[chat_task_id_value], queue="gen_chatapi_create", countdown=0) except Exception as exc: await _mark_dispatch_failed_and_raise(db, current_user=current_user, project_id=project_id_value, step_id=step_id_value, message=f"视频生成任务投递失败: {exc}") diff --git a/video-gen-api/app/config.py b/video-gen-api/app/config.py index 066c5528..f126cfcf 100644 --- a/video-gen-api/app/config.py +++ b/video-gen-api/app/config.py @@ -145,6 +145,17 @@ class Settings(BaseSettings): CELERY_STARTUP_RECOVERY_LOCK_KEY: str = "vg:celery:startup_recovery_lock" CELERY_STARTUP_RECOVERY_LOCK_TTL_SECONDS: int = 120 + # 模块异步任务容灾配置。 + # 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。 + # 不新增 worker 队列:恢复扫描仍走 gen_result_download,真实业务任务回到原始队列。 + MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 100 + MODULE_ASYNC_QUEUE_TIMEOUT_SECONDS: int = 5 * 60 + MODULE_ASYNC_LEASE_SECONDS: int = 10 * 60 + MODULE_ASYNC_LOCK_TTL_SECONDS: int = 10 * 60 + MODULE_ASYNC_REQUEUE_DELAY_SECONDS: int = 10 + MODULE_ASYNC_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:module_async:active" + MODULE_ASYNC_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:module_async:active_index" + MODULE_ASYNC_LOCK_KEY_PREFIX: str = "vg:celery:module_async:lock" RESOURCE_SIGN_SECRET: str = "resource-signature-secret-key-for-API-authentication" RESOURCE_SIGN_EXPIRE_SECONDS: int = 60 diff --git a/video-gen-api/app/services/module_async_recovery_service.py b/video-gen-api/app/services/module_async_recovery_service.py new file mode 100644 index 00000000..cff60b3c --- /dev/null +++ b/video-gen-api/app/services/module_async_recovery_service.py @@ -0,0 +1,615 @@ +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 _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=_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=_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: + await postpone_active_task( + object_type=object_type, + object_id=object_id, + delay_seconds=_lease_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} diff --git a/video-gen-api/app/services/shot_replicate_recovery_service.py b/video-gen-api/app/services/shot_replicate_recovery_service.py index 3470f89b..cbc6cd00 100644 --- a/video-gen-api/app/services/shot_replicate_recovery_service.py +++ b/video-gen-api/app/services/shot_replicate_recovery_service.py @@ -11,6 +11,7 @@ from app.enums.shot_replicate import ShotSplitStatusEnum from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary +from app.services.module_async_recovery_service import register_shot_split_task from app.tasks.celery_app import celery_app @@ -81,6 +82,7 @@ async def recover_one_split_segment(db: AsyncSession, segment: ShotReplicateSegm await refresh_task_set_split_summary(db, segment.task_set_id) await db.commit() + await register_shot_split_task(segment.id, task_set_id=segment.task_set_id) if celery_app: split_one_segment.apply_async( args=[segment.id], diff --git a/video-gen-api/app/tasks/__init__.py b/video-gen-api/app/tasks/__init__.py index 03e34c2c..81a62333 100644 --- a/video-gen-api/app/tasks/__init__.py +++ b/video-gen-api/app/tasks/__init__.py @@ -4,6 +4,10 @@ Celery autodiscover imports ``app.tasks``; importing task modules here ensures custom named tasks are registered when workers start. """ +import logging + +logger = logging.getLogger("video_gen") + try: from app.tasks import ( # noqa: F401 generation_create_tasks, @@ -12,7 +16,9 @@ try: generation_recovery_tasks, hot_opening_replicate_tasks, shot_replicate_tasks, - shot_replicate_flow_tasks + shot_replicate_flow_tasks, + module_async_recovery_tasks, ) except Exception: - pass + logger.exception("Celery 任务模块导入失败,worker 可能出现 unregistered task。") + raise diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 2f4e0f38..e203f2c0 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -59,6 +59,7 @@ if broker_url: "shot_replicate.recover_split_tasks_once": {"queue": "gen_result_download"}, "generation.recover_download_tasks_once": {"queue": "gen_result_download"}, "generation.recover_generation_tasks_once": {"queue": "gen_result_download"}, + "module_async.recover_module_async_tasks_once": {"queue": "gen_result_download"}, "app.tasks.cleanup.*": {"queue": "default"}, }, ) @@ -109,6 +110,7 @@ def on_worker_ready(sender=None, **kwargs): recover_generation_tasks_once, ) from app.tasks.shot_replicate_tasks import recover_split_tasks_once + from app.tasks.module_async_recovery_tasks import recover_module_async_tasks_once_task countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0)) @@ -127,6 +129,11 @@ def on_worker_ready(sender=None, **kwargs): queue="gen_result_download", priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, ) + recover_module_async_tasks_once_task.apply_async( + countdown=countdown + 15, + queue="gen_result_download", + priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + ) logger.info("启动容灾恢复任务已投递。countdown=%s", countdown) except Exception: diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index 08eea6e7..460e50f3 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -1,10 +1,11 @@ from app.tasks.async_runner import run_async import json -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Any, Optional from sqlalchemy import select +from app.config import settings from app.models.base import async_session from app.models.chat_generation_task import ChatGenerationTask from app.services.error_codes import extract_error_message @@ -16,6 +17,8 @@ from app.tasks.celery_app import celery_app ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} +def _now() -> datetime: + return datetime.now(timezone.utc) def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]: """ @@ -242,8 +245,13 @@ async def _run(task_id: str): await enqueue_download_task(db, task, reason="create_result_ready") else: - from app.tasks.generation_poll_tasks import poll_generation_task + from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active + await register_poll_active( + task, + reason="create_provider_success", + check_at=_now() + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)), + ) poll_generation_task.apply_async( args=[task.id], queue="gen_provider_poll", diff --git a/video-gen-api/app/tasks/hot_opening_replicate_tasks.py b/video-gen-api/app/tasks/hot_opening_replicate_tasks.py index a0228467..0dc9f5ba 100644 --- a/video-gen-api/app/tasks/hot_opening_replicate_tasks.py +++ b/video-gen-api/app/tasks/hot_opening_replicate_tasks.py @@ -1,45 +1,95 @@ from __future__ import annotations +from typing import Any + +from app.enums.hot_opening_replicate import HotOpeningStepCodeEnum, ModuleCodeEnum from app.models.base import async_session from app.services.hot_opening_replicate_service import run_image_prompt_optimize, run_video_prompt_optimize +from app.services.module_async_recovery_service import ( + OBJECT_MODULE_STEP, + TASK_HOT_IMAGE_PROMPT, + TASK_HOT_VIDEO_PROMPT, + acquire_object_lock, + cleanup_active_if_terminal, + mark_active_started, + release_object_lock, + register_module_step_task, +) from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app +MODULE = ModuleCodeEnum.HOT_OPENING_REPLICATE.value + async def _run_image_prompt(project_id: str, step_id: str | None = None): - async with async_session() as db: - await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) - await db.commit() + lock_token: str | None = None + if step_id: + lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) + if not lock_token: + # 同一个 step 的重复消息正在被其它 worker 处理:当前消息直接忽略,active registry 保留等待真正任务完成或 lease 过期。 + return None + await register_module_step_task( + module=MODULE, + project_id=project_id, + step_id=step_id, + step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + task_name=TASK_HOT_IMAGE_PROMPT, + ) + await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) + try: + async with async_session() as db: + result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) + await db.commit() + if step_id: + await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) + return result + finally: + if step_id: + await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) async def _run_video_prompt(project_id: str, step_id: str | None = None): - async with async_session() as db: - await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) - await db.commit() + lock_token: str | None = None + if step_id: + lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) + if not lock_token: + return None + await register_module_step_task( + module=MODULE, + project_id=project_id, + step_id=step_id, + step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + task_name=TASK_HOT_VIDEO_PROMPT, + ) + await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) + try: + async with async_session() as db: + result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) + await db.commit() + if step_id: + await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) + return result + finally: + if step_id: + await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) if celery_app: @celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None): - """手动触发后的图片 AI 提词任务。 - - 该任务路由到现有 gen_chatapi_create 队列,不需要新增 hot_opening worker。 - """ + """手动触发后的图片 AI 提词任务。""" return run_async(_run_image_prompt(project_id, step_id)) @celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None): - """手动触发后的视频 AI 提词任务。 - - 该任务路由到现有 gen_chatapi_create 队列,不需要新增 hot_opening worker。 - """ + """手动触发后的视频 AI 提词任务。""" return run_async(_run_video_prompt(project_id, step_id)) else: class _DisabledTask: - def delay(self, *args, **kwargs): + def delay(self, *args: Any, **kwargs: Any): raise RuntimeError("Celery is disabled") - def apply_async(self, *args, **kwargs): + def apply_async(self, *args: Any, **kwargs: Any): raise RuntimeError("Celery is disabled") start_image_prompt_optimize = _DisabledTask() diff --git a/video-gen-api/app/tasks/module_async_recovery_tasks.py b/video-gen-api/app/tasks/module_async_recovery_tasks.py new file mode 100644 index 00000000..6025eaec --- /dev/null +++ b/video-gen-api/app/tasks/module_async_recovery_tasks.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +from typing import Any + +from app.models.base import async_session +from app.services.module_async_recovery_service import recover_module_async_tasks_once +from app.tasks.async_runner import run_async +from app.tasks.celery_app import celery_app + + +async def _run_recover_module_async_tasks_once() -> dict[str, Any]: + async with async_session() as db: + return await recover_module_async_tasks_once(db) + + +if celery_app: + + @celery_app.task(name="module_async.recover_module_async_tasks_once") + def recover_module_async_tasks_once_task() -> dict[str, Any]: + return run_async(_run_recover_module_async_tasks_once()) + +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") + + recover_module_async_tasks_once_task = _DisabledTask() diff --git a/video-gen-api/app/tasks/shot_replicate_flow_tasks.py b/video-gen-api/app/tasks/shot_replicate_flow_tasks.py index 2d574458..ed603dcf 100644 --- a/video-gen-api/app/tasks/shot_replicate_flow_tasks.py +++ b/video-gen-api/app/tasks/shot_replicate_flow_tasks.py @@ -1,21 +1,76 @@ from __future__ import annotations +from typing import Any + +from app.enums.shot_replicate import ModuleCodeEnum, ShotReplicateStepCodeEnum from app.models.base import async_session +from app.services.module_async_recovery_service import ( + OBJECT_MODULE_STEP, + TASK_SHOT_IMAGE_PROMPT, + TASK_SHOT_VIDEO_PROMPT, + acquire_object_lock, + cleanup_active_if_terminal, + mark_active_started, + release_object_lock, + register_module_step_task, +) from app.services.shot_replicate_flow_service import run_image_prompt_optimize, run_video_prompt_optimize from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app +MODULE = ModuleCodeEnum.SHOT_REPLICATE.value + async def _run_image_prompt(project_id: str, step_id: str | None = None): - async with async_session() as db: - await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) - await db.commit() + lock_token: str | None = None + if step_id: + lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) + if not lock_token: + return None + await register_module_step_task( + module=MODULE, + project_id=project_id, + step_id=step_id, + step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + task_name=TASK_SHOT_IMAGE_PROMPT, + ) + await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) + try: + async with async_session() as db: + result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) + await db.commit() + if step_id: + await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) + return result + finally: + if step_id: + await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) async def _run_video_prompt(project_id: str, step_id: str | None = None): - async with async_session() as db: - await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) - await db.commit() + lock_token: str | None = None + if step_id: + lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) + if not lock_token: + return None + await register_module_step_task( + module=MODULE, + project_id=project_id, + step_id=step_id, + step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + task_name=TASK_SHOT_VIDEO_PROMPT, + ) + await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) + try: + async with async_session() as db: + result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) + await db.commit() + if step_id: + await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) + return result + finally: + if step_id: + await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) if celery_app: @@ -32,10 +87,10 @@ if celery_app: else: class _DisabledTask: - def delay(self, *args, **kwargs): + def delay(self, *args: Any, **kwargs: Any): raise RuntimeError("Celery is disabled") - def apply_async(self, *args, **kwargs): + def apply_async(self, *args: Any, **kwargs: Any): raise RuntimeError("Celery is disabled") start_image_prompt_optimize = _DisabledTask() diff --git a/video-gen-api/app/tasks/shot_replicate_tasks.py b/video-gen-api/app/tasks/shot_replicate_tasks.py index 6b98b29e..2bf43753 100644 --- a/video-gen-api/app/tasks/shot_replicate_tasks.py +++ b/video-gen-api/app/tasks/shot_replicate_tasks.py @@ -20,6 +20,20 @@ from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event from app.services.redis_registry_service import redis_acquire_lock, redis_release_lock +from app.services.module_async_recovery_service import ( + OBJECT_SHOT_SEGMENT_ANALYSIS, + OBJECT_SHOT_SPLIT_SEGMENT, + OBJECT_SHOT_TASK_SET_ANALYSIS, + acquire_object_lock, + cleanup_active_if_terminal, + mark_active_started, + postpone_active_task, + register_shot_segment_analysis_task, + register_shot_split_task, + register_shot_task_set_analysis_task, + release_object_lock, + remove_active_task, +) from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary from app.services.shot_video_analysis_service import analyze_video_for_shot_split from app.services.shot_video_split_service import split_video_segment_async @@ -65,6 +79,12 @@ async def _release_split_semaphore(lock_key: str | None, segment_id: str) -> Non async def _run_analyze_original_video(task_set_id: str) -> None: + lock_token = await acquire_object_lock(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) + if not lock_token: + return + await register_shot_task_set_analysis_task(task_set_id) + await mark_active_started(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) + task_set_user_id: str | None = None video_url: str | None = None try: @@ -77,8 +97,10 @@ async def _run_analyze_original_video(task_set_id: str) -> None: ) task_set = result.scalar_one_or_none() if not task_set: + await remove_active_task(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) return if task_set.analysis_status == ShotAnalysisStatusEnum.COMPLETED.value: + await remove_active_task(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) return task_set_user_id = task_set.user_id video_url = task_set.video_url @@ -107,6 +129,7 @@ async def _run_analyze_original_video(task_set_id: str) -> None: task_set = result.scalar_one_or_none() if not task_set: await db.rollback() + await remove_active_task(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) return result_json = analyzed.result task_set.original_video_content = str(result_json.get("原视频内容") or "无") @@ -119,6 +142,7 @@ async def _run_analyze_original_video(task_set_id: str) -> None: task_set.status = ShotTaskSetStatusEnum.ANALYSIS_COMPLETED.value task_set.analysis_error_message = None await db.commit() + await cleanup_active_if_terminal(db, object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) log_module_prompt_event( event_type="SHOT_ANALYSIS_SUCCESS", @@ -154,6 +178,9 @@ async def _run_analyze_original_video(task_set_id: str) -> None: task_set.analysis_status = ShotAnalysisStatusEnum.FAILED.value task_set.analysis_error_message = str(exc) await db.commit() + await cleanup_active_if_terminal(db, object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) + else: + await remove_active_task(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id) log_module_error( module=MODULE, event_type="SHOT_ANALYSIS_FAILED", @@ -163,9 +190,17 @@ async def _run_analyze_original_video(task_set_id: str) -> None: detail={"task_set_id": task_set_id, "video_url": video_url, "analysis_mode": "full_breakdown"}, exc=exc, ) + finally: + await release_object_lock(object_type=OBJECT_SHOT_TASK_SET_ANALYSIS, object_id=task_set_id, token=lock_token) async def _run_analyze_custom_segment_video(segment_id: str) -> None: + lock_token = await acquire_object_lock(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) + if not lock_token: + return + await register_shot_segment_analysis_task(segment_id) + await mark_active_started(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) + user_id: str | None = None task_set_id: str | None = None video_url: str | None = None @@ -179,12 +214,15 @@ async def _run_analyze_custom_segment_video(segment_id: str) -> None: ) segment = result.scalar_one_or_none() if not segment or not segment.segment_video_url: + await remove_active_task(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) return if segment.analysis_status == ShotSegmentAnalysisStatusEnum.COMPLETED.value: + await remove_active_task(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) return user_id = segment.user_id task_set_id = segment.task_set_id video_url = segment.segment_video_url + await register_shot_segment_analysis_task(segment_id, task_set_id=task_set_id) segment.analysis_status = ShotSegmentAnalysisStatusEnum.PROCESSING.value segment.analysis_error_message = None await db.commit() @@ -210,6 +248,7 @@ async def _run_analyze_custom_segment_video(segment_id: str) -> None: segment = result.scalar_one_or_none() if not segment: await db.rollback() + await remove_active_task(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) return result_json = analyzed.result segment.original_video_content = str(result_json.get("原视频内容") or "无") @@ -222,10 +261,11 @@ async def _run_analyze_custom_segment_video(segment_id: str) -> None: segment.analysis_status = ShotSegmentAnalysisStatusEnum.COMPLETED.value segment.analysis_error_message = None await db.commit() + await cleanup_active_if_terminal(db, object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) log_module_prompt_event( event_type="SHOT_SEGMENT_ANALYSIS_SUCCESS", - project_id=task_set_id or segment_id, + project_id=task_set_id, step_id=segment_id, user_id=user_id or "", module=MODULE, @@ -258,6 +298,9 @@ async def _run_analyze_custom_segment_video(segment_id: str) -> None: segment.analysis_status = ShotSegmentAnalysisStatusEnum.FAILED.value segment.analysis_error_message = str(exc) await db.commit() + await cleanup_active_if_terminal(db, object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) + else: + await remove_active_task(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id) log_module_error( module=MODULE, event_type="SHOT_SEGMENT_ANALYSIS_FAILED", @@ -268,6 +311,8 @@ async def _run_analyze_custom_segment_video(segment_id: str) -> None: detail={"segment_id": segment_id, "task_set_id": task_set_id, "video_url": video_url, "analysis_mode": "summary_only"}, exc=exc, ) + finally: + await release_object_lock(object_type=OBJECT_SHOT_SEGMENT_ANALYSIS, object_id=segment_id, token=lock_token) async def _run_split_one_segment(segment_id: str) -> None: @@ -280,6 +325,9 @@ async def _run_split_one_segment(segment_id: str) -> None: if not segment_lock_token: return + await register_shot_split_task(segment_id) + await mark_active_started(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) + semaphore_key: str | None = None user_id: str | None = None task_set_id: str | None = None @@ -294,8 +342,15 @@ async def _run_split_one_segment(segment_id: str) -> None: message="拆镜 ffmpeg 并发闸门已满,稍后重试", detail={"segment_id": segment_id, "reason": "semaphore_full"}, ) + delay = max(1, int(settings.MODULE_ASYNC_REQUEUE_DELAY_SECONDS or 10)) + await postpone_active_task( + object_type=OBJECT_SHOT_SPLIT_SEGMENT, + object_id=segment_id, + delay_seconds=delay, + reason="semaphore_full", + ) if celery_app: - split_one_segment.apply_async(args=[segment_id], queue=SPLIT_QUEUE, countdown=10, priority=settings.DOWNLOAD_TASK_PRIORITY_NORMAL) + split_one_segment.apply_async(args=[segment_id], queue=SPLIT_QUEUE, countdown=delay, priority=settings.DOWNLOAD_TASK_PRIORITY_NORMAL) return async with async_session() as db: @@ -307,6 +362,7 @@ async def _run_split_one_segment(segment_id: str) -> None: ) segment = result.scalar_one_or_none() if not segment: + await remove_active_task(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) return user_id = segment.user_id task_set_id = segment.task_set_id @@ -318,8 +374,10 @@ async def _run_split_one_segment(segment_id: str) -> None: ) task_set = task_set_result.scalar_one_or_none() if not task_set: + await remove_active_task(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) return if segment.split_status == ShotSplitStatusEnum.COMPLETED.value and segment.segment_video_url: + await remove_active_task(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) return validate_split_range( @@ -338,6 +396,8 @@ async def _run_split_one_segment(segment_id: str) -> None: task_set.status = ShotTaskSetStatusEnum.SPLITTING.value task_set.split_status = ShotSplitStatusEnum.PROCESSING.value await db.commit() + await register_shot_split_task(segment_id, task_set_id=segment.task_set_id) + await mark_active_started(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) source_path = task_set.video_path date_dir = (segment.created_at or now).strftime("%Y/%m/%d") @@ -379,6 +439,7 @@ async def _run_split_one_segment(segment_id: str) -> None: ) segment = result.scalar_one_or_none() if not segment: + await remove_active_task(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) return segment.segment_video_url = split_result.url segment.segment_video_path = split_result.path @@ -389,6 +450,7 @@ async def _run_split_one_segment(segment_id: str) -> None: segment.split_last_error = None await refresh_task_set_split_summary(db, segment.task_set_id) await db.commit() + await cleanup_active_if_terminal(db, object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) log_module_event_file( module=MODULE, @@ -407,6 +469,7 @@ async def _run_split_one_segment(segment_id: str) -> None: ) if segment.source_mode == ShotSegmentSourceModeEnum.CUSTOM.value and celery_app: + await register_shot_segment_analysis_task(segment.id, task_set_id=segment.task_set_id) analyze_custom_segment_video.apply_async(args=[segment.id], queue=ANALYSIS_QUEUE, countdown=0) except Exception as exc: @@ -421,6 +484,7 @@ async def _run_split_one_segment(segment_id: str) -> None: ) segment = result.scalar_one_or_none() if not segment: + await remove_active_task(object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) return user_id = user_id or segment.user_id task_set_id = task_set_id or segment.task_set_id @@ -437,9 +501,18 @@ async def _run_split_one_segment(segment_id: str) -> None: await refresh_task_set_split_summary(db, segment.task_set_id) await db.commit() - if segment.split_status == ShotSplitStatusEnum.RETRY_WAITING.value and celery_app: + if segment.split_status == ShotSplitStatusEnum.RETRY_WAITING.value: next_retry_delay = max(1, int(((segment.split_next_retry_at or _now()) - _now()).total_seconds())) - split_one_segment.apply_async(args=[segment_id], queue=SPLIT_QUEUE, countdown=next_retry_delay, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER) + await postpone_active_task( + object_type=OBJECT_SHOT_SPLIT_SEGMENT, + object_id=segment_id, + delay_seconds=next_retry_delay, + reason="split_retry_waiting", + ) + if celery_app: + split_one_segment.apply_async(args=[segment_id], queue=SPLIT_QUEUE, countdown=next_retry_delay, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER) + elif final_failed: + await cleanup_active_if_terminal(db, object_type=OBJECT_SHOT_SPLIT_SEGMENT, object_id=segment_id) log_module_error( module=MODULE,