Files
video-gen/video-gen-api/app/tasks/shot_replicate_flow_tasks.py
T

98 lines
3.7 KiB
Python

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):
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):
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:
@celery_app.task(name="shot_replicate.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):
return run_async(_run_image_prompt(project_id, step_id))
@celery_app.task(name="shot_replicate.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):
return run_async(_run_video_prompt(project_id, step_id))
else:
class _DisabledTask:
def delay(self, *args: Any, **kwargs: Any):
raise RuntimeError("Celery is disabled")
def apply_async(self, *args: Any, **kwargs: Any):
raise RuntimeError("Celery is disabled")
start_image_prompt_optimize = _DisabledTask()
start_video_prompt_optimize = _DisabledTask()