557 lines
20 KiB
Python
557 lines
20 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
import uuid
|
|
from datetime import datetime, timedelta, timezone
|
|
from typing import Any, Iterable
|
|
|
|
from celery import current_task
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.config import settings
|
|
from app.enums.common import ModuleStepStatusEnum
|
|
from app.enums.celery_queue import CeleryQueue
|
|
from app.enums.celery_runtime import CeleryRuntimeDomain
|
|
from app.enums.hot_opening_replicate import HotOpeningStepCodeEnum, ModuleCodeEnum as HotModuleCodeEnum
|
|
from app.enums.shot_replicate import (
|
|
ModuleCodeEnum as ShotModuleCodeEnum,
|
|
ShotReplicateStepCodeEnum,
|
|
)
|
|
from app.models.module_generation_project import ModuleGenerationProject
|
|
from app.models.module_generation_step import ModuleGenerationStep
|
|
from app.services.redis_registry_service import (
|
|
datetime_to_epoch,
|
|
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.services.celery_runtime.runtime_service import CeleryRuntimeLease, RuntimeIdentity, runtime_lock_values
|
|
from app.tasks.celery_app import celery_app
|
|
|
|
logger = logging.getLogger("video_gen")
|
|
|
|
QUEUE_CREATE = CeleryQueue.GEN_CHATAPI_CREATE.value
|
|
|
|
OBJECT_MODULE_STEP = "module_step"
|
|
|
|
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_MODULE_V2_VIDEO_PROMPT = "module_generation_v2.start_video_prompt_optimize"
|
|
|
|
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,
|
|
}
|
|
|
|
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,
|
|
)
|
|
existing = await redis_get_registry_payloads(
|
|
hash_key=_hash_key(),
|
|
item_ids=[item_id],
|
|
log_context="module_async_active",
|
|
)
|
|
if item_id in existing:
|
|
merged = dict(existing[item_id])
|
|
merged.update(payload)
|
|
payload = merged
|
|
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 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,
|
|
)
|
|
|
|
|
|
_OBJECT_LEASES: dict[str, CeleryRuntimeLease] = {}
|
|
|
|
|
|
def _current_task_metadata() -> tuple[str, str, str | None]:
|
|
task = current_task
|
|
task_name = str(getattr(task, "name", "") or "module_async.unknown")
|
|
request = getattr(task, "request", None)
|
|
delivery = getattr(request, "delivery_info", None) or {}
|
|
queue = str(delivery.get("routing_key") or delivery.get("exchange") or QUEUE_CREATE)
|
|
celery_task_id = str(getattr(request, "id", "") or "") or None
|
|
return task_name, queue, celery_task_id
|
|
|
|
|
|
async def acquire_object_lock(*, object_type: str, object_id: str) -> str | None:
|
|
task_name, queue, celery_task_id = _current_task_metadata()
|
|
token = uuid.uuid4().hex
|
|
lease = await CeleryRuntimeLease.acquire(
|
|
identity=RuntimeIdentity(
|
|
domain=CeleryRuntimeDomain.MODULE_ASYNC.value,
|
|
owner_type=object_type,
|
|
owner_id=object_id,
|
|
attempt_no=1,
|
|
task_name=task_name,
|
|
queue=queue,
|
|
registry_item_id=_item_id(object_type, object_id),
|
|
),
|
|
token=token,
|
|
lock_key=_lock_key(object_type, object_id),
|
|
hash_key=_hash_key(),
|
|
zset_key=_zset_key(),
|
|
ttl_seconds=max(60, int(settings.MODULE_ASYNC_LOCK_TTL_SECONDS or _lease_seconds())),
|
|
heartbeat_interval_seconds=max(10, min(30, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30))),
|
|
pipeline_stage="processing",
|
|
extra_payload={
|
|
"object_type": object_type,
|
|
"object_id": object_id,
|
|
"celery_task_id": celery_task_id,
|
|
},
|
|
)
|
|
if lease is None:
|
|
return None
|
|
_OBJECT_LEASES[token] = lease
|
|
return token
|
|
|
|
|
|
async def ensure_object_lock_owned(*, token: str | None) -> None:
|
|
if not token:
|
|
raise RuntimeError("module async execution token is missing")
|
|
lease = _OBJECT_LEASES.get(token)
|
|
if lease is None:
|
|
raise RuntimeError("module async execution lease is unavailable")
|
|
await lease.ensure_owned()
|
|
|
|
|
|
async def release_object_lock(*, object_type: str, object_id: str, token: str | None) -> None:
|
|
if not token:
|
|
return
|
|
lease = _OBJECT_LEASES.pop(token, None)
|
|
if lease is not None:
|
|
await lease.close()
|
|
return
|
|
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:
|
|
await remove_active_task(object_type=object_type, object_id=object_id)
|
|
return True
|
|
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
|
|
|
|
|
|
def _payload_args(payload: dict[str, Any]) -> list[Any]:
|
|
args = payload.get("args")
|
|
if isinstance(args, list):
|
|
return args
|
|
if str(payload.get("object_type") or "") == OBJECT_MODULE_STEP:
|
|
return [payload.get("project_id"), payload.get("step_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_MODULE_STEP:
|
|
await remove_active_task(object_type=object_type, object_id=object_id)
|
|
return "remove_legacy_non_module_payload"
|
|
|
|
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",
|
|
)
|
|
valid_payloads: dict[str, dict[str, Any]] = {}
|
|
step_ids: set[str] = set()
|
|
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"
|
|
)
|
|
results["remove_missing_payload"] = results.get("remove_missing_payload", 0) + 1
|
|
continue
|
|
object_type = str(payload.get("object_type") or "")
|
|
object_id = str(payload.get("object_id") or "")
|
|
if object_type != OBJECT_MODULE_STEP or not object_id:
|
|
await redis_remove_registry_item(
|
|
hash_key=_hash_key(), zset_key=_zset_key(), item_id=item_id, log_context="module_async_active"
|
|
)
|
|
results["remove_legacy_non_module_payload"] = results.get("remove_legacy_non_module_payload", 0) + 1
|
|
continue
|
|
valid_payloads[item_id] = payload
|
|
step_ids.add(object_id)
|
|
|
|
step_map: dict[str, ModuleGenerationStep] = {}
|
|
if step_ids:
|
|
step_result = await db.execute(
|
|
select(ModuleGenerationStep).where(ModuleGenerationStep.id.in_(step_ids))
|
|
)
|
|
step_map = {str(step.id): step for step in step_result.scalars().all()}
|
|
lock_keys = [_lock_key(OBJECT_MODULE_STEP, step_id) for step_id in step_ids]
|
|
lock_values = await runtime_lock_values(lock_keys)
|
|
|
|
for item_id, payload in valid_payloads.items():
|
|
object_id = str(payload.get("object_id") or "")
|
|
step = step_map.get(object_id)
|
|
if not step or step.deleted_at is not None or step.status in TERMINAL_STEP_STATUSES:
|
|
await remove_active_task(object_type=OBJECT_MODULE_STEP, object_id=object_id)
|
|
results["remove_terminal"] = results.get("remove_terminal", 0) + 1
|
|
continue
|
|
if lock_values.get(_lock_key(OBJECT_MODULE_STEP, object_id)):
|
|
await postpone_active_task(
|
|
object_type=OBJECT_MODULE_STEP,
|
|
object_id=object_id,
|
|
delay_seconds=_lease_seconds(),
|
|
reason="live_runtime_lock",
|
|
)
|
|
results["skip_live_step_lock"] = results.get("skip_live_step_lock", 0) + 1
|
|
continue
|
|
action = await _recover_payload_from_redis(db, item_id, payload)
|
|
results[action] = results.get(action, 0) + 1
|
|
return results
|
|
|
|
|
|
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())
|
|
project_flow_map: dict[str, str] = {}
|
|
project_ids = list({step.project_id for step in steps})
|
|
if project_ids:
|
|
project_result = await db.execute(
|
|
select(ModuleGenerationProject.id, ModuleGenerationProject.flow_version).where(
|
|
ModuleGenerationProject.id.in_(project_ids),
|
|
ModuleGenerationProject.deleted_at.is_(None),
|
|
)
|
|
)
|
|
project_flow_map = {
|
|
str(project_id): str(flow_version or "v1")
|
|
for project_id, flow_version in project_result.all()
|
|
}
|
|
results: dict[str, int] = {}
|
|
dispatches: list[tuple[str, str, str, str, str]] = []
|
|
lock_keys = [_lock_key(OBJECT_MODULE_STEP, str(step.id)) for step in steps]
|
|
live_locks = await runtime_lock_values(lock_keys)
|
|
for step in steps:
|
|
if live_locks.get(_lock_key(OBJECT_MODULE_STEP, str(step.id))):
|
|
results["skip_live_step_lock"] = results.get("skip_live_step_lock", 0) + 1
|
|
continue
|
|
if project_flow_map.get(step.project_id, "v1") == "v2" and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
|
|
task_name = TASK_MODULE_V2_VIDEO_PROMPT
|
|
elif 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
|
|
|
|
dispatches.append((task_name, step.module, step.project_id, step.id, step.step_code))
|
|
await db.commit()
|
|
for task_name, module, project_id, step_id, step_code in dispatches:
|
|
await register_module_step_task(
|
|
module=module,
|
|
project_id=project_id,
|
|
step_id=step_id,
|
|
step_code=step_code,
|
|
task_name=task_name,
|
|
queue=QUEUE_CREATE,
|
|
)
|
|
try:
|
|
_send_task(task_name, args=[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
|
|
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]:
|
|
"""恢复爆款开头与拆镜复刻的短 LLM 提词步骤。
|
|
|
|
长视频分析与 FFmpeg 切片分别由独立队列和恢复服务负责,
|
|
避免多套恢复链路重复投递同一业务任务。
|
|
"""
|
|
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())
|
|
|
|
return {"checked": sum(results.values()), "results": results}
|