from __future__ import annotations import logging import uuid from datetime import datetime, timedelta, timezone from typing import Any from sqlalchemy import select from app.config import settings from app.enums.celery_queue import CeleryQueue, CeleryTaskName from app.enums.celery_runtime import CeleryRuntimeDomain from app.enums.private_portrait import ( PrivatePortraitAssetStatus, PrivatePortraitEventSource, PrivatePortraitEventType, PrivatePortraitRemoteDeleteStatus, ) from app.models import async_session from app.models.virtual_portrait_v3 import VpV3Asset, VpV3Project from app.services.celery_runtime.recovery_service import guard_periodic_recovery from app.services.celery_runtime.runtime_service import CeleryRuntimeLease, RuntimeIdentity from app.services.operation_log_service import log_operation_error, log_operation_event from app.services.redis_registry_service import RedisExecutionLockLease from app.services.virtual_portrait_v3.asset_service import ( V3_DOMAIN, delete_v3_asset_remote, sync_asset_status, ) from app.services.virtual_portrait_v3.project_service import ( V3_DOMAIN as V3_PROJECT_DOMAIN, delete_v3_project_remote, ) from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app logger = logging.getLogger(__name__) QUEUE = CeleryQueue.GEN_PRIVATE_PORTRAIT.value # 北京时间(UTC+8)统一基准:V3 业务所有时间写入 / 时间比较唯一参考 _BJ_TZ = timezone(timedelta(hours=8)) def _bj_now() -> datetime: """返回当前北京时间(UTC+8)naive datetime(去掉 tzinfo,与 DB naive 存储一致)。""" return datetime.now(_BJ_TZ).replace(tzinfo=None) def _now() -> datetime: """统一使用北京时间基准,与业务写入保持一致。""" return _bj_now() def _naive(dt: datetime | None) -> datetime | None: """把 datetime 统一成 naive 北京时间(去掉 tzinfo),避免 offset-aware vs naive 比较报错。 DB 列是 DateTime(timezone=True) 但业务写入都是北京时间(naive), 读回时根据方言可能变成 aware 或仍为 naive,比较前统一去掉 tzinfo。 """ if dt is None: return None return dt.replace(tzinfo=None) if dt.tzinfo is not None else dt def _retry_countdown(retries: int) -> int: return min(300, 30 * (2 ** max(0, retries))) async def _rollback_and_reraise( db, *, event_type: str, exc: BaseException, detail: dict[str, Any] | None = None, **kwargs: Any, ): await db.rollback() log_operation_error( domain=V3_DOMAIN, event_type=event_type, source=PrivatePortraitEventSource.CELERY.value, exc=exc, detail=detail, **kwargs, ) raise exc async def _acquire_v3_runtime( *, domain: str, owner_type: str, owner_id: str, task_name: str, hash_key: str, zset_key: str, lock_prefix: str, ) -> CeleryRuntimeLease | None: token = uuid.uuid4().hex return await CeleryRuntimeLease.acquire( identity=RuntimeIdentity( domain=domain, owner_type=owner_type, owner_id=owner_id, attempt_no=1, task_name=task_name, queue=QUEUE, ), lock_key=f"{lock_prefix}:{owner_type}:{owner_id}:attempt:1", hash_key=hash_key, zset_key=zset_key, token=token, ttl_seconds=max(60, int(settings.VP_V3_RUNTIME_LOCK_TTL_SECONDS or 180)), heartbeat_interval_seconds=max(10, int(settings.VP_V3_RUNTIME_HEARTBEAT_SECONDS or 30)), pipeline_stage="processing", ) # --------------------------------------------------------------------------- # 轮询:单条素材 # --------------------------------------------------------------------------- async def _run_poll_v3_asset(asset_id: str) -> None: lease = await _acquire_v3_runtime( domain=CeleryRuntimeDomain.PRIVATE_PORTRAIT_POLL.value, owner_type="asset", owner_id=asset_id, task_name=CeleryTaskName.VP_V3_POLL_ASSET.value, hash_key=settings.VP_V3_POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.VP_V3_POLL_ACTIVE_REDIS_ZSET_KEY, lock_prefix=settings.VP_V3_POLL_LOCK_KEY_PREFIX, ) if lease is None: logger.info("vp_v3 poll asset skip: runtime lease not acquired (asset_id=%s)", asset_id) return try: async with async_session() as db: try: row = (await db.execute( select(VpV3Asset.api_key_id).where(VpV3Asset.remote_asset_id == asset_id).limit(1) )).scalar_one_or_none() if not row: return asset = await sync_asset_status( db, api_key_id=str(row), asset_id=asset_id, execution_guard=lease.ensure_owned, ) await lease.ensure_owned() await db.commit() logger.info( "vp_v3 poll asset synced: asset_id=%s status=%s poll_count=%s", asset_id, asset.status, asset.poll_count, ) except Exception as exc: logger.exception("vp_v3 poll asset failed: %s", asset_id) await _rollback_and_reraise( db, event_type=PrivatePortraitEventType.ASSET_POLL_FAILED.value, exc=exc, asset_id=asset_id, detail={"celery_task": CeleryTaskName.VP_V3_POLL_ASSET.value}, ) finally: await lease.close() # --------------------------------------------------------------------------- # 轮询:每分钟批量扫描到期素材并分发轮询任务 # --------------------------------------------------------------------------- async def _dispatch_v3_due_assets() -> int: barrier = await guard_periodic_recovery() if barrier is not None: logger.info("vp_v3 dispatch due assets skip: periodic recovery barrier active") return 0 lock = await RedisExecutionLockLease.acquire( lock_key=settings.VP_V3_DISPATCH_LOCK_KEY, ttl_seconds=55, renew_interval_seconds=20, log_context="vp_v3_poll_dispatch", ) if lock is None: logger.info("vp_v3 dispatch due assets skip: dispatch lock not acquired") return 0 async with lock: async with async_session() as db: now_naive = _naive(_now()) rows = await db.execute( select(VpV3Asset) .where( VpV3Asset.deleted_at.is_(None), VpV3Asset.status == PrivatePortraitAssetStatus.CREATING.value, VpV3Asset.next_poll_at.is_not(None), ) .order_by(VpV3Asset.next_poll_at.asc(), VpV3Asset.id.asc()) .limit(settings.VP_V3_ASSET_POLL_BATCH_SIZE or 50) .with_for_update(skip_locked=True) ) assets = list(rows.scalars().all()) # next_poll_at <= now 在内存里过滤(统一 naive 比较,避免 aware vs naive 报错) assets = [a for a in assets if _naive(a.next_poll_at) is not None and _naive(a.next_poll_at) <= now_naive] dispatches: list[tuple[str, int]] = [] queue_hold_until_naive = now_naive + timedelta(seconds=120) for asset in assets: poll_no = int(asset.poll_count or 0) + 1 dispatches.append((str(asset.remote_asset_id), poll_no)) asset.next_poll_at = queue_hold_until_naive await db.commit() for asset_id, poll_no in dispatches: poll_v3_asset_status.apply_async( args=[asset_id], queue=QUEUE, countdown=0, task_id=f"vp-v3-poll:{asset_id}:attempt:{poll_no}", ) log_operation_event( domain=V3_DOMAIN, event_type=PrivatePortraitEventType.SYNC_DUE_ASSETS_DONE.value, event_status="success", source=PrivatePortraitEventSource.CELERY.value, detail={"matched_count": len(dispatches), "dispatched_count": len(dispatches)}, ) return len(dispatches) # --------------------------------------------------------------------------- # 删除:素材 / 项目远端删除(已存在) # --------------------------------------------------------------------------- async def _run_delete_v3_asset(asset_id: str) -> None: """执行 V3 素材远端删除。""" lease = await _acquire_v3_runtime( domain=CeleryRuntimeDomain.PRIVATE_PORTRAIT_DELETE.value, owner_type="asset", owner_id=asset_id, task_name=CeleryTaskName.VP_V3_DELETE_ASSET.value, hash_key=settings.VP_V3_DELETE_ACTIVE_REDIS_HASH_KEY, zset_key=settings.VP_V3_DELETE_ACTIVE_REDIS_ZSET_KEY, lock_prefix=settings.VP_V3_DELETE_LOCK_KEY_PREFIX, ) if lease is None: logger.info("vp_v3 delete asset skip: runtime lease not acquired (asset_id=%s)", asset_id) return try: async with async_session() as db: try: await delete_v3_asset_remote( db, asset_id=asset_id, execution_guard=lease.ensure_owned, ) await lease.ensure_owned() await db.commit() except Exception as exc: await _rollback_and_reraise( db, event_type=PrivatePortraitEventType.ASSET_DELETE_REMOTE_FAILED.value, exc=exc, detail={"asset_id": asset_id}, ) finally: await lease.close() async def _run_delete_v3_project(project_id: str) -> int: """执行 V3 项目远端删除(级联删除素材 + 项目)。""" lease = await _acquire_v3_runtime( domain=CeleryRuntimeDomain.PRIVATE_PORTRAIT_DELETE.value, owner_type="project", owner_id=project_id, task_name=CeleryTaskName.VP_V3_DELETE_PROJECT.value, hash_key=settings.VP_V3_DELETE_ACTIVE_REDIS_HASH_KEY, zset_key=settings.VP_V3_DELETE_ACTIVE_REDIS_ZSET_KEY, lock_prefix=settings.VP_V3_DELETE_LOCK_KEY_PREFIX, ) if lease is None: logger.info("vp_v3 delete project skip: runtime lease not acquired (project_id=%s)", project_id) return 0 try: async with async_session() as db: try: await delete_v3_project_remote( db, project_id=project_id, ) await lease.ensure_owned() await db.commit() except Exception as exc: await _rollback_and_reraise( db, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_FAILED.value, exc=exc, detail={"project_id": project_id}, ) finally: await lease.close() return 0 # --------------------------------------------------------------------------- # 删除恢复:每 5 分钟扫描 pending/failed 的 project/asset 再投递 # --------------------------------------------------------------------------- async def _dispatch_v3_remote_delete_recovery() -> dict[str, int]: barrier = await guard_periodic_recovery() if barrier is not None: logger.info("vp_v3 delete recovery skip: periodic recovery barrier active") return {"asset_count": 0, "project_count": 0, "total_count": 0} lock = await RedisExecutionLockLease.acquire( lock_key=settings.VP_V3_DELETE_RECOVERY_LOCK_KEY, ttl_seconds=240, renew_interval_seconds=30, log_context="vp_v3_delete_recovery", ) if lock is None: logger.info("vp_v3 delete recovery skip: recovery lock not acquired") return {"asset_count": 0, "project_count": 0, "total_count": 0} async with lock: statuses = [ PrivatePortraitRemoteDeleteStatus.PENDING.value, PrivatePortraitRemoteDeleteStatus.FAILED.value, ] batch_size = max(1, int(settings.VP_V3_REMOTE_DELETE_RECOVERY_BATCH_SIZE or 50)) async with async_session() as db: asset_rows = await db.execute( select(VpV3Asset.id) .where(VpV3Asset.remote_delete_status.in_(statuses)) .order_by(VpV3Asset.updated_at.asc(), VpV3Asset.id.asc()) .limit(batch_size) ) asset_ids = [str(value) for value in asset_rows.scalars().all()] remaining = max(0, batch_size - len(asset_ids)) project_ids: list[str] = [] if remaining: project_rows = await db.execute( select(VpV3Project.id) .where(VpV3Project.remote_delete_status.in_(statuses)) .order_by(VpV3Project.updated_at.asc(), VpV3Project.id.asc()) .limit(remaining) ) project_ids = [str(value) for value in project_rows.scalars().all()] await db.rollback() for asset_id in asset_ids: delete_v3_asset_remote_task.apply_async( args=[asset_id], queue=QUEUE, task_id=f"vp-v3-delete-asset:{asset_id}" ) for project_id in project_ids: delete_v3_project_remote_task.apply_async( args=[project_id], queue=QUEUE, task_id=f"vp-v3-delete-project:{project_id}" ) return { "asset_count": len(asset_ids), "project_count": len(project_ids), "total_count": len(asset_ids) + len(project_ids), } # --------------------------------------------------------------------------- # Celery 任务注册 # --------------------------------------------------------------------------- @celery_app.task( name=CeleryTaskName.VP_V3_POLL_ASSET.value, queue=QUEUE, bind=True, max_retries=5, default_retry_delay=30, ) def poll_v3_asset_status(self, asset_id: str) -> None: """V3 素材单条状态轮询(Celery 任务)。""" logger.info("vp_v3 poll task START: asset_id=%s task_id=%s", asset_id, self.request.id) try: return run_async(_run_poll_v3_asset(asset_id)) except Exception as exc: raise self.retry(exc=exc, countdown=_retry_countdown(self.request.retries)) @celery_app.task( name=CeleryTaskName.VP_V3_SYNC_DUE_ASSETS.value, queue=QUEUE, bind=True, max_retries=3, default_retry_delay=60, ) def sync_v3_due_assets(self) -> int: """每分钟扫描 V3 到期素材并分发轮询任务(beat schedule)。""" logger.info("vp_v3 sync_due_assets START: task_id=%s", self.request.id) try: return run_async(_dispatch_v3_due_assets()) except Exception as exc: raise self.retry(exc=exc, countdown=_retry_countdown(self.request.retries)) @celery_app.task( name=CeleryTaskName.VP_V3_DELETE_ASSET.value, queue=QUEUE, bind=True, max_retries=3, default_retry_delay=60, ) def delete_v3_asset_remote_task(self, asset_id: str) -> None: """V3 素材远端删除 Celery 任务。""" logger.info("vp_v3 delete asset START: asset_id=%s task_id=%s", asset_id, self.request.id) try: return run_async(_run_delete_v3_asset(asset_id)) except Exception as exc: raise self.retry(exc=exc, countdown=_retry_countdown(self.request.retries)) @celery_app.task( name=CeleryTaskName.VP_V3_DELETE_PROJECT.value, queue=QUEUE, bind=True, max_retries=3, default_retry_delay=60, ) def delete_v3_project_remote_task(self, project_id: str) -> int: """V3 项目远端删除 Celery 任务。""" logger.info("vp_v3 delete project START: project_id=%s task_id=%s", project_id, self.request.id) try: return run_async(_run_delete_v3_project(project_id)) except Exception as exc: raise self.retry(exc=exc, countdown=_retry_countdown(self.request.retries)) @celery_app.task( name=CeleryTaskName.VP_V3_RECOVER_REMOTE_DELETES.value, queue=QUEUE, bind=True, max_retries=3, default_retry_delay=60, ) def recover_v3_remote_deletes(self) -> dict[str, int]: """每 5 分钟扫描 V3 pending/failed 远端删除记录并重新投递(beat schedule)。""" logger.info("vp_v3 recover_remote_deletes START: task_id=%s", self.request.id) try: return run_async(_dispatch_v3_remote_delete_recovery()) except Exception as exc: raise self.retry(exc=exc, countdown=_retry_countdown(self.request.retries))