450 lines
16 KiB
Python
450 lines
16 KiB
Python
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))
|