Files
video-gen/video-gen-api/app/tasks/vp_v3_asset_tasks.py
T
root 0c511f3451 1、增加调用 AI 视频生成能力和虚拟素材库管理的对外api
2、增加后台apikkey管理
3、增加apikey单独的模型定价
4、增加apikey调用情况
5、完善所有数据的注释增加
2026-08-06 13:13:28 +08:00

450 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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+8naive 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))