Files
video-gen/video-gen-api/app/services/module_async_recovery_service.py
T
2026-07-24 09:18:05 +08:00

708 lines
27 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 ModuleProjectStatusEnum, ModuleStepStatusEnum
from app.enums.celery_queue import CeleryQueue
from app.enums.celery_runtime import CeleryRuntimeDomain
from app.enums.credit_record import (
CreditRecordBillingScene,
CreditRecordChargeKind,
CreditRecordOwnerType,
)
from app.enums.llm_billing import LlmBillingConfigKey, LlmBillingLedgerState
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.services.llm_billing import LlmBillingContext, LlmHoldValidation, get_llm_ledger_states
from app.services.llm_billing.config import get_llm_billing_policy
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 _step_llm_billing_context(step: ModuleGenerationStep) -> LlmBillingContext:
is_image_prompt = step.step_code in {
HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
}
if step.module == HOT_MODULE:
billing_scene = (
CreditRecordBillingScene.HOT_OPENING_IMAGE_PROMPT_OPTIMIZE.value
if is_image_prompt
else CreditRecordBillingScene.HOT_OPENING_VIDEO_PROMPT_OPTIMIZE.value
)
else:
billing_scene = (
CreditRecordBillingScene.SHOT_IMAGE_PROMPT_OPTIMIZE.value
if is_image_prompt
else CreditRecordBillingScene.SHOT_VIDEO_PROMPT_OPTIMIZE.value
)
return LlmBillingContext(
user_id=str(step.user_id),
owner_type=CreditRecordOwnerType.MODULE_GENERATION_STEP.value,
owner_id=str(step.id),
attempt_no=max(1, int(step.version or 1)),
charge_kind=CreditRecordChargeKind.TEXT_PROMPT.value,
billing_scene=billing_scene,
source_module=str(step.module),
source_project_id=str(step.project_id),
source_step_id=str(step.id),
source_step_code=str(step.step_code),
related_id=str(step.id),
hold_config_key=(
LlmBillingConfigKey.HOLD_MODULE_IMAGE_PROMPT.value
if is_image_prompt
else LlmBillingConfigKey.HOLD_MODULE_VIDEO_PROMPT.value
),
description_prefix="模块AI提词优化",
trace_id=f"module-recovery:{step.id}:attempt:{max(1, int(step.version or 1))}",
)
async def _load_step_billing_validations(
db: AsyncSession,
steps: Iterable[ModuleGenerationStep],
) -> dict[str, LlmHoldValidation]:
step_list = list(steps)
if not step_list:
return {}
contexts = {str(step.id): _step_llm_billing_context(step) for step in step_list}
policies = {}
for config_key in {ctx.hold_config_key for ctx in contexts.values() if ctx.hold_config_key}:
policies[config_key] = await get_llm_billing_policy(db, config_key=config_key)
# 无论当前配置是否关闭,都批量读取历史 attempt 流水:运行中的 active HOLD
# 必须继续结算,不能因后台关闭计费而被当成 bypass 遗留冻结。
ledger_states = await get_llm_ledger_states(db, contexts.values())
output: dict[str, LlmHoldValidation] = {}
for step_id, ctx in contexts.items():
policy = policies.get(ctx.hold_config_key)
ledger = ledger_states.get(
ctx.hold_biz_key,
LlmHoldValidation(False, 0.0, LlmBillingLedgerState.MISSING, "ledger_not_loaded"),
)
if ledger.state == LlmBillingLedgerState.ACTIVE:
output[step_id] = ledger
elif ledger.state == LlmBillingLedgerState.MISSING and policy is not None and policy.bypassed:
output[step_id] = LlmHoldValidation(
True,
0.0,
LlmBillingLedgerState.BILLING_BYPASSED,
"billing_disabled",
)
elif ledger.state == LlmBillingLedgerState.MISSING and (policy is None or not policy.valid):
output[step_id] = LlmHoldValidation(
False,
0.0,
LlmBillingLedgerState.INVALID,
"billing_config_invalid",
)
else:
output[step_id] = ledger
return output
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 has_live_object_lock(*, object_type: str, object_id: str) -> bool:
"""判断对象是否已被 worker 领取,供投递异常补偿规避不确定投递竞态。"""
lock_key = _lock_key(object_type, object_id)
values = await runtime_lock_values([lock_key])
return bool(values.get(lock_key))
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 _module_task_id(task_name: str, args: list[Any]) -> str | None:
step_id = str(args[-1]) if args else ""
if not step_id:
return None
if task_name == TASK_HOT_IMAGE_PROMPT:
return f"hot-opening:image-prompt:{step_id}"
if task_name == TASK_HOT_VIDEO_PROMPT:
return f"hot-opening:video-prompt:{step_id}"
if task_name == TASK_SHOT_IMAGE_PROMPT:
return f"shot-replicate:image-prompt:{step_id}"
if task_name == TASK_SHOT_VIDEO_PROMPT:
return f"shot-replicate:video-prompt:{step_id}"
if task_name == TASK_MODULE_V2_VIDEO_PROMPT:
return f"module-v2-video-prompt:{step_id}"
return f"module-async:{task_name}:{step_id}"
def _send_task(task_name: str, *, args: list[Any], queue: str, countdown: int = 0, priority: int | None = None) -> None:
if celery_app is None:
raise RuntimeError("Celery 未启用,不能恢复投递模块 LLM 任务")
if not task_name or not queue:
raise ValueError("恢复投递缺少 task_name 或 queue")
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,
task_id=_module_task_id(task_name, args),
)
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] = {}
project_map: dict[str, ModuleGenerationProject] = {}
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()}
project_ids = {str(step.project_id) for step in step_map.values()}
if project_ids:
project_result = await db.execute(
select(ModuleGenerationProject).where(ModuleGenerationProject.id.in_(project_ids))
)
project_map = {str(project.id): project for project in project_result.scalars().all()}
billing_validations = await _load_step_billing_validations(db, step_map.values())
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
billing_validation = billing_validations.get(object_id)
if billing_validation is None or not billing_validation.can_execute:
state = billing_validation.state.value if billing_validation else LlmBillingLedgerState.MISSING.value
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = f"LLM账务状态异常({state}),恢复任务已终止"
step.completed_at = _now()
project = project_map.get(str(step.project_id))
if project is not None:
project.status = ModuleProjectStatusEnum.FAILED.value
project.error_message = step.error_message
await remove_active_task(object_type=OBJECT_MODULE_STEP, object_id=object_id)
result_key = f"remove_billing_{state}"
results[result_key] = results.get(result_key, 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_map: dict[str, ModuleGenerationProject] = {}
project_ids = list({step.project_id for step in steps})
if project_ids:
project_result = await db.execute(
select(ModuleGenerationProject).where(
ModuleGenerationProject.id.in_(project_ids),
ModuleGenerationProject.deleted_at.is_(None),
)
)
project_map = {str(project.id): project for project in project_result.scalars().all()}
project_flow_map = {
project_id: str(project.flow_version or "v1")
for project_id, project in project_map.items()
}
billing_validations = await _load_step_billing_validations(db, steps)
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
billing_validation = billing_validations.get(str(step.id))
if billing_validation is None or not billing_validation.can_execute:
state = billing_validation.state.value if billing_validation else LlmBillingLedgerState.MISSING.value
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = f"LLM账务状态异常({state}),恢复任务已终止"
step.completed_at = _now()
project = project_map.get(str(step.project_id))
if project is not None:
project.status = ModuleProjectStatusEnum.FAILED.value
project.error_message = step.error_message
result_key = f"db_step_billing_{state}"
results[result_key] = results.get(result_key, 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}