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 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, LlmChargeValidation, get_llm_ledger_states 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), 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, LlmChargeValidation]: step_list = list(steps) if not step_list: return {} contexts = {str(step.id): _step_llm_billing_context(step) for step in step_list} # 场景积分消费已经在 API 层完成;恢复任务只认同一业务 attempt 的账务执行记录, # 不再读取旧 SystemConfig 冻结配置;缺少场景积分消费记录时直接终止恢复。 ledger_states = await get_llm_ledger_states(db, contexts.values()) output: dict[str, LlmChargeValidation] = {} for step_id, ctx in contexts.items(): ledger = ledger_states.get( ctx.charge_biz_key, LlmChargeValidation(False, 0.0, LlmBillingLedgerState.MISSING, "ledger_not_loaded"), ) if ledger.state == LlmBillingLedgerState.ACTIVE: output[step_id] = ledger elif ledger.state == LlmBillingLedgerState.MISSING: output[step_id] = LlmChargeValidation( 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}