Files
video-gen/video-gen-api/app/services/private_portrait/asset_service.py
T
2026-07-07 15:18:52 +08:00

723 lines
42 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 json
from datetime import datetime, timedelta, timezone
from typing import Any
from urllib.parse import urlencode
from fastapi import HTTPException
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.private_portrait import (
PRIVATE_PORTRAIT_ASSET_POLL_INTERVAL_SECONDS,
PRIVATE_PORTRAIT_ASSET_POLL_MAX_COUNT,
PRIVATE_PORTRAIT_ASSET_URI_PREFIX,
PRIVATE_PORTRAIT_ENABLED_ASSET_TYPES,
PRIVATE_PORTRAIT_REAL_PERSON_GROUP_TYPE,
PRIVATE_PORTRAIT_SUCCESS_RESULT_CODE,
PRIVATE_PORTRAIT_VALIDATE_TOKEN_EXPIRE_MINUTES,
PRIVATE_PORTRAIT_VIDEO_ASSET_POLL_INTERVAL_SECONDS,
PRIVATE_PORTRAIT_VIDEO_ASSET_POLL_MAX_COUNT,
PrivatePortraitAssetGroupStatus,
PrivatePortraitAssetStatus,
PrivatePortraitAssetType,
PrivatePortraitEventSource,
PrivatePortraitEventStatus,
PrivatePortraitEventType,
PrivatePortraitLibraryType,
PrivatePortraitProjectStatus,
PrivatePortraitRemoteDeleteStatus,
PrivatePortraitValidateSessionStatus,
)
from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitAssetGroup, PrivatePortraitProject, PrivatePortraitValidateSession
from app.schemas.private_portrait import PrivatePortraitAssetCreate, PrivatePortraitAssetOut, PrivatePortraitSelectableAssetOut, PrivatePortraitValidateSessionOut
from app.services.operation_log_service import log_operation_error, log_operation_event
from app.services.private_portrait.ark_client import ArkPrivateAssetClient
from app.services.private_portrait.project_service import get_user_project, refresh_project_counters
from app.services.private_portrait.quota_service import (
count_user_counting_assets,
ensure_private_portrait_asset_quota_available,
get_user_private_portrait_config,
set_user_private_portrait_limit,
)
from app.utils.id_gen import generate_id
DOMAIN = "private_portrait"
PRIVATE_PORTRAIT_VIDEO_MIN_DURATION_SECONDS = 2
PRIVATE_PORTRAIT_VIDEO_MAX_DURATION_SECONDS = 15
def _json(data: Any) -> str | None:
if data is None:
return None
return json.dumps(data, ensure_ascii=False, default=str)
def _loads(data: str | None) -> Any:
if not data:
return None
try:
return json.loads(data)
except Exception:
return None
def _exception_message(exc: Exception) -> str:
if isinstance(exc, HTTPException):
detail = exc.detail
if isinstance(detail, dict):
message = detail.get("message") or detail.get("detail") or detail
return str(message)
return str(detail)
return str(exc)
def _public_url(url: str) -> str:
if url.startswith(("http://", "https://", PRIVATE_PORTRAIT_ASSET_URI_PREFIX)):
return url
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
def _callback_url(session_id: str, callback_redirect_url: str | None = None) -> str:
base = f"{settings.BASE_URL.rstrip('/')}/api/private-portrait/validate-callback"
params = {"session_id": session_id}
if callback_redirect_url:
params["redirect_url"] = callback_redirect_url
return f"{base}?{urlencode(params)}"
def _remote_group_name(user_id: str, project_name: str) -> str:
safe_name = "".join(ch if ch.isalnum() or ch in "-_" else "_" for ch in project_name.strip())[:80]
return f"{user_id}-{safe_name}"[:128]
def _asset_display_url(asset: PrivatePortraitAsset) -> str | None:
return asset.preview_url or asset.remote_url or asset.source_url or None
def _provider_url(asset: PrivatePortraitAsset) -> str | None:
return f"{PRIVATE_PORTRAIT_ASSET_URI_PREFIX}{asset.remote_asset_id}" if asset.remote_asset_id else None
def _poll_interval_seconds(asset_type: str) -> int:
if asset_type == PrivatePortraitAssetType.VIDEO.value:
return PRIVATE_PORTRAIT_VIDEO_ASSET_POLL_INTERVAL_SECONDS
return PRIVATE_PORTRAIT_ASSET_POLL_INTERVAL_SECONDS
def _poll_max_count(asset_type: str) -> int:
if asset_type == PrivatePortraitAssetType.VIDEO.value:
return PRIVATE_PORTRAIT_VIDEO_ASSET_POLL_MAX_COUNT
return PRIVATE_PORTRAIT_ASSET_POLL_MAX_COUNT
def _assert_enabled_asset_type(asset_type: str) -> None:
if asset_type not in {item.value for item in PrivatePortraitAssetType}:
raise HTTPException(status_code=400, detail="asset_type 仅支持 Image/VideoAudio 暂未开放")
if asset_type not in PRIVATE_PORTRAIT_ENABLED_ASSET_TYPES:
raise HTTPException(status_code=400, detail="Audio 暂未开放,当前仅支持 Image/Video")
def _assert_private_asset_video_duration(payload: PrivatePortraitAssetCreate) -> None:
if payload.asset_type != PrivatePortraitAssetType.VIDEO.value:
return
duration = payload.video_duration
if duration is None:
raise HTTPException(status_code=400, detail="Video 素材必须提供 video_duration")
if duration < PRIVATE_PORTRAIT_VIDEO_MIN_DURATION_SECONDS:
raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 {PRIVATE_PORTRAIT_VIDEO_MIN_DURATION_SECONDS} 秒")
if duration > PRIVATE_PORTRAIT_VIDEO_MAX_DURATION_SECONDS:
raise HTTPException(status_code=400, detail=f"视频素材最长不能超过 {PRIVATE_PORTRAIT_VIDEO_MAX_DURATION_SECONDS} 秒")
def validate_session_to_out(session: PrivatePortraitValidateSession, *, include_user: bool = False) -> PrivatePortraitValidateSessionOut:
return PrivatePortraitValidateSessionOut(
id=session.id,
user_id=session.user_id if include_user else None,
project_id=session.project_id,
byted_token=session.byted_token,
h5_link=session.h5_link,
callback_url=session.callback_url,
result_code=session.result_code,
algorithm_base_resp_code=session.algorithm_base_resp_code,
verify_type=session.verify_type,
status=session.status,
remote_group_id=session.remote_group_id,
remote_project_name=session.remote_project_name,
expired_at=session.expired_at,
error_message=session.error_message,
created_at=session.created_at,
updated_at=session.updated_at,
)
def asset_to_out(asset: PrivatePortraitAsset, *, project_name: str | None = None, include_user: bool = False) -> PrivatePortraitAssetOut:
display_url = _asset_display_url(asset)
return PrivatePortraitAssetOut(
id=asset.id,
user_id=asset.user_id if include_user else None,
project_id=asset.project_id,
project_name=project_name,
group_id=asset.group_id,
library_type=asset.library_type,
remote_group_id=asset.remote_group_id,
remote_asset_id=asset.remote_asset_id,
remote_project_name=asset.remote_project_name,
asset_type=asset.asset_type,
name=asset.name,
source_url=asset.source_url,
preview_url=asset.preview_url,
display_url=display_url,
provider_url=_provider_url(asset),
remote_url=asset.remote_url,
remote_url_expired_at=asset.remote_url_expired_at,
video_duration=asset.video_duration,
video_cover_url=asset.video_cover_url,
file_size=asset.file_size,
mime_type=asset.mime_type,
status=asset.status,
moderation=_loads(asset.moderation_json),
last_poll_at=asset.last_poll_at,
next_poll_at=asset.next_poll_at,
poll_count=asset.poll_count or 0,
remote_delete_status=asset.remote_delete_status,
remote_deleted_at=asset.remote_deleted_at,
remote_delete_error=asset.remote_delete_error,
error_message=asset.error_message,
created_at=asset.created_at,
updated_at=asset.updated_at,
)
async def _get_existing_active_group(db: AsyncSession, *, project_id: str, library_type: str | None = None) -> PrivatePortraitAssetGroup | None:
filters = [
PrivatePortraitAssetGroup.project_id == project_id,
PrivatePortraitAssetGroup.status == PrivatePortraitAssetGroupStatus.ACTIVE.value,
PrivatePortraitAssetGroup.deleted_at.is_(None),
]
if library_type:
filters.append(PrivatePortraitAssetGroup.library_type == library_type)
return (
await db.execute(
select(PrivatePortraitAssetGroup)
.where(*filters)
.order_by(PrivatePortraitAssetGroup.created_at.desc())
.limit(1)
)
).scalar_one_or_none()
async def _ensure_project_can_validate(db: AsyncSession, *, project: PrivatePortraitProject) -> PrivatePortraitValidateSession | None:
if project.library_type != PrivatePortraitLibraryType.REAL_PERSON.value:
raise HTTPException(status_code=400, detail="虚拟人像项目不支持真人认证")
active_group = await _get_existing_active_group(db, project_id=project.id, library_type=project.library_type)
if active_group or project.status == PrivatePortraitProjectStatus.ACTIVE.value:
raise HTTPException(status_code=409, detail="该真人素材项目已完成认证,不能重复认证")
success_session = (
await db.execute(
select(PrivatePortraitValidateSession)
.where(
PrivatePortraitValidateSession.project_id == project.id,
PrivatePortraitValidateSession.status == PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value,
)
.order_by(PrivatePortraitValidateSession.updated_at.desc())
.limit(1)
)
).scalar_one_or_none()
if success_session:
raise HTTPException(status_code=409, detail="该真人素材项目已完成认证,不能重复认证")
now = datetime.now(timezone.utc)
pending_session = (
await db.execute(
select(PrivatePortraitValidateSession)
.where(
PrivatePortraitValidateSession.project_id == project.id,
PrivatePortraitValidateSession.status.in_([PrivatePortraitValidateSessionStatus.CREATED.value, PrivatePortraitValidateSessionStatus.CALLBACK_SUCCESS.value]),
PrivatePortraitValidateSession.expired_at.is_not(None),
PrivatePortraitValidateSession.expired_at > now,
)
.order_by(PrivatePortraitValidateSession.created_at.desc())
.limit(1)
)
).scalar_one_or_none()
return pending_session
async def create_validate_session(db: AsyncSession, *, user_id: str, project_id: str, callback_redirect_url: str | None = None) -> PrivatePortraitValidateSession:
project = await get_user_project(db, user_id=user_id, project_id=project_id, library_type=PrivatePortraitLibraryType.REAL_PERSON.value)
reusable_session = await _ensure_project_can_validate(db, project=project)
if reusable_session:
return reusable_session
project.status = PrivatePortraitProjectStatus.VALIDATING.value
session = PrivatePortraitValidateSession(
id=generate_id(),
user_id=user_id,
project_id=project.id,
status=PrivatePortraitValidateSessionStatus.CREATED.value,
remote_project_name=project.remote_project_name,
expired_at=datetime.now(timezone.utc) + timedelta(minutes=PRIVATE_PORTRAIT_VALIDATE_TOKEN_EXPIRE_MINUTES),
)
session.callback_url = _callback_url(session.id, callback_redirect_url)
db.add(session)
await db.flush()
try:
resp = await ArkPrivateAssetClient().create_visual_validate_session(project_name=project.remote_project_name, callback_url=session.callback_url)
session.byted_token = resp.get("BytedToken") or resp.get("bytedToken")
session.h5_link = resp.get("H5Link") or resp.get("h5Link")
session.raw_response_json = _json(resp)
await db.flush()
await db.refresh(session)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_SESSION_CREATE.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, session_id=session.id, detail={"remote_project_name": project.remote_project_name, "library_type": project.library_type})
return session
except Exception as exc:
session.status = PrivatePortraitValidateSessionStatus.FAILED.value
session.error_message = _exception_message(exc)
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
await db.flush()
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_SESSION_CREATE_FAILED.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, session_id=session.id, exc=exc)
raise
async def get_validate_session(db: AsyncSession, *, user_id: str | None, session_id: str) -> PrivatePortraitValidateSession:
filters = [PrivatePortraitValidateSession.id == session_id]
if user_id is not None:
filters.append(PrivatePortraitValidateSession.user_id == user_id)
session = (await db.execute(select(PrivatePortraitValidateSession).where(*filters).limit(1))).scalar_one_or_none()
if not session:
raise HTTPException(status_code=404, detail="真人认证会话不存在")
return session
async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_params: dict[str, Any]) -> PrivatePortraitValidateSession:
session = await get_validate_session(db, user_id=None, session_id=session_id)
if session.status == PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value:
return session
session.raw_callback_json = _json(query_params)
session.result_code = str(query_params.get("resultCode") or query_params.get("result_code") or "") or None
session.algorithm_base_resp_code = str(query_params.get("algorithmBaseRespCode") or query_params.get("algorithm_base_resp_code") or "") or None
session.verify_type = str(query_params.get("verify_type") or query_params.get("verifyType") or "") or None
token = query_params.get("bytedToken") or query_params.get("byted_token") or session.byted_token
if token:
session.byted_token = str(token)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_CALLBACK_RECEIVED.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, detail={"query_params": query_params, "remote_project_name": session.remote_project_name})
project = (await db.execute(select(PrivatePortraitProject).where(PrivatePortraitProject.id == session.project_id).limit(1))).scalar_one_or_none()
if session.result_code != PRIVATE_PORTRAIT_SUCCESS_RESULT_CODE:
session.status = PrivatePortraitValidateSessionStatus.CALLBACK_FAILED.value
session.error_message = f"真人认证失败:resultCode={session.result_code}"
if project and project.status != PrivatePortraitProjectStatus.ACTIVE.value:
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
await db.flush()
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_CALLBACK_FAILED.value, event_status=PrivatePortraitEventStatus.FAILED.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, error=session.error_message)
return session
session.status = PrivatePortraitValidateSessionStatus.CALLBACK_SUCCESS.value
if not session.byted_token:
session.status = PrivatePortraitValidateSessionStatus.FAILED.value
session.error_message = "Callback 未返回 BytedToken"
if project and project.status != PrivatePortraitProjectStatus.ACTIVE.value:
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
await db.flush()
raise HTTPException(status_code=400, detail=session.error_message)
try:
existing_group = await _get_existing_active_group(db, project_id=session.project_id, library_type=PrivatePortraitLibraryType.REAL_PERSON.value)
if existing_group:
session.remote_group_id = existing_group.remote_group_id
session.status = PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value
if project:
project.status = PrivatePortraitProjectStatus.ACTIVE.value
await db.flush()
await db.refresh(session)
return session
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_GET_RESULT_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, detail={"remote_project_name": session.remote_project_name})
resp = await ArkPrivateAssetClient().get_visual_validate_result(project_name=session.remote_project_name, byted_token=session.byted_token)
group_id = resp.get("GroupId") or resp.get("groupId")
if not group_id:
raise RuntimeError("GetVisualValidateResult 未返回 GroupId")
session.remote_group_id = group_id
session.status = PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value
session.raw_response_json = _json(resp)
if not project:
raise RuntimeError("真人素材项目不存在")
remote_group_name = _remote_group_name(session.user_id, project.name)
group = PrivatePortraitAssetGroup(
id=generate_id(),
user_id=session.user_id,
project_id=session.project_id,
library_type=PrivatePortraitLibraryType.REAL_PERSON.value,
remote_group_id=group_id,
remote_group_name=remote_group_name,
remote_project_name=session.remote_project_name,
group_type=PRIVATE_PORTRAIT_REAL_PERSON_GROUP_TYPE,
status=PrivatePortraitAssetGroupStatus.ACTIVE.value,
raw_response_json=_json(resp),
)
db.add(group)
project.status = PrivatePortraitProjectStatus.ACTIVE.value
await db.flush()
try:
await ArkPrivateAssetClient().update_asset_group(project_name=session.remote_project_name, group_id=group_id, name=remote_group_name, title=remote_group_name, description=project.description)
except Exception as exc:
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_GROUP_UPDATE_REMOTE_FAILED.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, group_id=group.id, exc=exc)
await refresh_project_counters(db, [session.project_id])
await db.flush()
await db.refresh(session)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_GET_RESULT_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, group_id=group.id, detail={"remote_group_id": group_id, "remote_project_name": session.remote_project_name, "library_type": PrivatePortraitLibraryType.REAL_PERSON.value})
return session
except Exception as exc:
session.status = PrivatePortraitValidateSessionStatus.FAILED.value
session.error_message = _exception_message(exc)
if project and project.status != PrivatePortraitProjectStatus.ACTIVE.value:
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
await db.flush()
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_GET_RESULT_FAILED.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, exc=exc)
raise
async def get_project_active_group(db: AsyncSession, *, user_id: str, project_id: str, library_type: str | None = None) -> PrivatePortraitAssetGroup:
project = await get_user_project(db, user_id=user_id, project_id=project_id, library_type=library_type)
if project.status != PrivatePortraitProjectStatus.ACTIVE.value:
detail = "请先完成真人授权认证,再上传素材" if project.library_type == PrivatePortraitLibraryType.REAL_PERSON.value else "虚拟人像素材组尚未创建成功,不能上传素材"
raise HTTPException(status_code=400, detail=detail)
result = await db.execute(
select(PrivatePortraitAssetGroup)
.where(
PrivatePortraitAssetGroup.user_id == user_id,
PrivatePortraitAssetGroup.project_id == project_id,
PrivatePortraitAssetGroup.library_type == project.library_type,
PrivatePortraitAssetGroup.status == PrivatePortraitAssetGroupStatus.ACTIVE.value,
PrivatePortraitAssetGroup.deleted_at.is_(None),
)
.order_by(PrivatePortraitAssetGroup.created_at.desc())
.limit(1)
)
group = result.scalar_one_or_none()
if not group:
raise HTTPException(status_code=400, detail="项目没有可用的远程素材组")
return group
async def create_asset(
db: AsyncSession,
*,
user_id: str,
project_id: str,
payload: PrivatePortraitAssetCreate,
library_type: str | None = None,
) -> PrivatePortraitAsset:
_assert_enabled_asset_type(payload.asset_type)
_assert_private_asset_video_duration(payload)
project = await get_user_project(db, user_id=user_id, project_id=project_id, library_type=library_type)
if project.status != PrivatePortraitProjectStatus.ACTIVE.value:
raise HTTPException(status_code=400, detail="项目未激活,不能上传素材")
limit, current_count = await ensure_private_portrait_asset_quota_available(db, user_id=user_id, project_id=project_id, library_type=project.library_type, asset_type=payload.asset_type)
group = await get_project_active_group(db, user_id=user_id, project_id=project.id, library_type=project.library_type)
public_url = _public_url(payload.url)
asset = PrivatePortraitAsset(
id=generate_id(),
user_id=user_id,
project_id=project.id,
group_id=group.id,
library_type=project.library_type,
remote_group_id=group.remote_group_id,
remote_project_name=project.remote_project_name,
asset_type=payload.asset_type,
name=payload.name,
source_url=public_url,
preview_url=payload.url,
video_duration=payload.video_duration,
video_cover_url=payload.video_cover_url,
file_size=payload.file_size,
mime_type=payload.mime_type,
status=PrivatePortraitAssetStatus.CREATING.value,
)
db.add(asset)
await db.flush()
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_CREATE_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, group_id=group.id, asset_id=asset.id, detail={"asset_limit": limit, "used_asset_count": current_count, "library_type": project.library_type, "asset_type": payload.asset_type, "remote_project_name": project.remote_project_name})
try:
remote_resp = await ArkPrivateAssetClient().create_asset(project_name=project.remote_project_name, group_id=group.remote_group_id, url=public_url, asset_type=payload.asset_type, name=payload.name)
remote_asset_id = remote_resp.get("Id") or remote_resp.get("AssetId") or remote_resp.get("assetId")
if not remote_asset_id:
raise RuntimeError("CreateAsset 未返回素材 ID")
now = datetime.now(timezone.utc)
asset.remote_asset_id = remote_asset_id
asset.status = PrivatePortraitAssetStatus.PROCESSING.value
asset.next_poll_at = now + timedelta(seconds=_poll_interval_seconds(asset.asset_type))
asset.raw_response_json = _json(remote_resp)
await refresh_project_counters(db, [project.id])
await db.flush()
await db.refresh(asset)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_CREATE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, group_id=group.id, asset_id=asset.id, detail={"remote_asset_id": remote_asset_id, "remote_project_name": project.remote_project_name, "library_type": project.library_type, "asset_type": asset.asset_type})
return asset
except Exception as exc:
asset.status = PrivatePortraitAssetStatus.FAILED.value
asset.error_message = _exception_message(exc)
await db.flush()
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_CREATE_FAILED.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, group_id=group.id, asset_id=asset.id, exc=exc)
raise
async def sync_asset_status(db: AsyncSession, *, user_id: str | None, asset_id: str) -> PrivatePortraitAsset:
filters = [PrivatePortraitAsset.id == asset_id]
if user_id is not None:
filters.append(PrivatePortraitAsset.user_id == user_id)
asset = (await db.execute(select(PrivatePortraitAsset).where(*filters).limit(1))).scalar_one_or_none()
if not asset:
raise HTTPException(status_code=404, detail="私域人像素材不存在")
if asset.deleted_at is not None:
raise HTTPException(status_code=400, detail="私域人像素材已删除")
if not asset.remote_asset_id:
raise HTTPException(status_code=400, detail="私域人像素材尚未创建远程 Asset")
source = PrivatePortraitEventSource.CELERY.value if user_id is None else PrivatePortraitEventSource.API.value
log_operation_event(
domain=DOMAIN,
event_type=PrivatePortraitEventType.ASSET_SYNC_START.value,
event_status=PrivatePortraitEventStatus.PENDING.value,
source=source,
user_id=asset.user_id,
project_id=asset.project_id,
asset_id=asset.id,
detail={"status": asset.status, "poll_count": int(asset.poll_count or 0), "remote_asset_id": asset.remote_asset_id, "remote_project_name": asset.remote_project_name, "library_type": asset.library_type, "asset_type": asset.asset_type},
)
try:
remote_resp = await ArkPrivateAssetClient(for_celery=(user_id is None)).get_asset(project_name=asset.remote_project_name, asset_id=asset.remote_asset_id)
status = remote_resp.get("Status") or remote_resp.get("status")
now = datetime.now(timezone.utc)
asset.last_poll_at = now
asset.poll_count = int(asset.poll_count or 0) + 1
asset.raw_response_json = _json(remote_resp)
if status:
asset.status = status
asset.remote_url = remote_resp.get("URL") or remote_resp.get("url") or asset.remote_url
asset.moderation_json = _json(remote_resp.get("Moderation") or remote_resp.get("moderation"))
max_count = _poll_max_count(asset.asset_type)
if asset.status == PrivatePortraitAssetStatus.PROCESSING.value and asset.poll_count >= max_count:
asset.status = PrivatePortraitAssetStatus.FAILED.value
asset.error_message = "素材入库轮询超时"
asset.next_poll_at = None
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_POLL_TIMEOUT.value, event_status=PrivatePortraitEventStatus.FAILED.value, source=source, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, detail={"poll_count": asset.poll_count, "max_count": max_count, "remote_asset_id": asset.remote_asset_id, "library_type": asset.library_type, "asset_type": asset.asset_type}, error=asset.error_message)
elif asset.status == PrivatePortraitAssetStatus.PROCESSING.value:
asset.next_poll_at = now + timedelta(seconds=_poll_interval_seconds(asset.asset_type))
else:
asset.next_poll_at = None
if asset.status == PrivatePortraitAssetStatus.FAILED.value and not asset.error_message:
asset.error_message = remote_resp.get("ErrorMessage") or remote_resp.get("error_message") or "素材入库失败"
await refresh_project_counters(db, [asset.project_id])
await db.flush()
await db.refresh(asset)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_SYNC_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=source, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, detail={"status": asset.status, "remote_asset_id": asset.remote_asset_id, "next_poll_at": asset.next_poll_at, "poll_count": asset.poll_count, "library_type": asset.library_type, "asset_type": asset.asset_type})
return asset
except Exception as exc:
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_SYNC_FAILED.value, source=source, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, exc=exc)
raise
async def list_assets(
db: AsyncSession,
*,
user_id: str | None,
project_id: str | None = None,
status: str | None = None,
keyword: str | None = None,
page: int = 1,
page_size: int = 20,
library_type: str | None = None,
asset_type: str | None = None,
) -> tuple[list[PrivatePortraitAsset], int, dict[str, str]]:
page = max(1, page)
page_size = min(max(1, page_size), 100)
filters = [PrivatePortraitAsset.deleted_at.is_(None)]
if user_id:
filters.append(PrivatePortraitAsset.user_id == user_id)
if project_id:
filters.append(PrivatePortraitAsset.project_id == project_id)
if library_type:
filters.append(PrivatePortraitAsset.library_type == library_type)
if asset_type:
filters.append(PrivatePortraitAsset.asset_type == asset_type)
if status:
filters.append(PrivatePortraitAsset.status == status)
if keyword:
filters.append(PrivatePortraitAsset.name.ilike(f"%{keyword.strip()}%"))
total = (await db.execute(select(func.count(PrivatePortraitAsset.id)).where(*filters))).scalar_one()
result = await db.execute(select(PrivatePortraitAsset).where(*filters).order_by(PrivatePortraitAsset.created_at.desc()).offset((page - 1) * page_size).limit(page_size))
assets = list(result.scalars().all())
project_ids = list({asset.project_id for asset in assets})
project_name_map: dict[str, str] = {}
if project_ids:
rows = await db.execute(select(PrivatePortraitProject.id, PrivatePortraitProject.name).where(PrivatePortraitProject.id.in_(project_ids)))
project_name_map = {pid: name for pid, name in rows.all()}
return assets, int(total or 0), project_name_map
async def list_selectable_assets(db: AsyncSession, *, user_id: str, project_id: str | None = None, keyword: str | None = None, page: int = 1, page_size: int = 20, library_type: str | None = None, asset_type: str | None = None) -> tuple[list[PrivatePortraitSelectableAssetOut], int]:
assets, total, project_name_map = await list_assets(db, user_id=user_id, project_id=project_id, status=PrivatePortraitAssetStatus.ACTIVE.value, keyword=keyword, page=page, page_size=page_size, library_type=library_type, asset_type=asset_type)
return [
PrivatePortraitSelectableAssetOut(
id=asset.id,
project_id=asset.project_id,
project_name=project_name_map.get(asset.project_id, ""),
library_type=asset.library_type,
name=asset.name,
asset_type=asset.asset_type,
preview_url=asset.preview_url or asset.remote_url,
display_url=_asset_display_url(asset),
provider_url=_provider_url(asset),
video_duration=asset.video_duration,
video_cover_url=asset.video_cover_url,
status=asset.status,
created_at=asset.created_at,
)
for asset in assets
], total
async def soft_delete_asset(db: AsyncSession, *, user_id: str, asset_id: str, library_type: str | None = None) -> PrivatePortraitAsset:
filters = [PrivatePortraitAsset.id == asset_id, PrivatePortraitAsset.user_id == user_id, PrivatePortraitAsset.deleted_at.is_(None)]
if library_type:
filters.append(PrivatePortraitAsset.library_type == library_type)
asset = (await db.execute(select(PrivatePortraitAsset).where(*filters).limit(1))).scalar_one_or_none()
if not asset:
raise HTTPException(status_code=404, detail="私域人像素材不存在")
now = datetime.now(timezone.utc)
asset.deleted_at = now
asset.status = PrivatePortraitAssetStatus.LOCAL_DELETED.value
asset.remote_delete_status = PrivatePortraitRemoteDeleteStatus.PENDING.value
await refresh_project_counters(db, [asset.project_id])
await db.flush()
await db.refresh(asset)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_DELETE_LOCAL.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=asset.project_id, asset_id=asset.id, detail={"remote_asset_id": asset.remote_asset_id, "remote_project_name": asset.remote_project_name, "library_type": asset.library_type, "asset_type": asset.asset_type})
return asset
async def delete_asset_remote(db: AsyncSession, *, asset_id: str) -> None:
asset = (await db.execute(select(PrivatePortraitAsset).where(PrivatePortraitAsset.id == asset_id).limit(1))).scalar_one_or_none()
if not asset:
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_DELETE_REMOTE_START.value, event_status=PrivatePortraitEventStatus.SKIPPED.value, source=PrivatePortraitEventSource.CELERY.value, asset_id=asset_id, message="远程删除跳过:本地素材不存在")
return
if not asset.remote_asset_id:
asset.remote_delete_status = PrivatePortraitRemoteDeleteStatus.SKIPPED.value
asset.remote_delete_error = None
await db.flush()
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_DELETE_REMOTE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SKIPPED.value, source=PrivatePortraitEventSource.CELERY.value, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, message="远程删除跳过:素材没有 remote_asset_id")
return
now = datetime.now(timezone.utc)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_DELETE_REMOTE_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CELERY.value, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, detail={"remote_asset_id": asset.remote_asset_id, "remote_project_name": asset.remote_project_name, "library_type": asset.library_type, "asset_type": asset.asset_type})
try:
await ArkPrivateAssetClient(for_celery=True).delete_asset(project_name=asset.remote_project_name, asset_id=asset.remote_asset_id)
asset.status = PrivatePortraitAssetStatus.REMOTE_DELETED.value
asset.remote_delete_status = PrivatePortraitRemoteDeleteStatus.SUCCESS.value
asset.remote_deleted_at = now
asset.remote_delete_error = None
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_DELETE_REMOTE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.CELERY.value, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, detail={"remote_asset_id": asset.remote_asset_id, "remote_project_name": asset.remote_project_name, "library_type": asset.library_type})
except Exception as exc:
asset.status = PrivatePortraitAssetStatus.DELETE_FAILED.value
asset.remote_delete_status = PrivatePortraitRemoteDeleteStatus.FAILED.value
asset.remote_delete_error = str(exc)
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_DELETE_REMOTE_FAILED.value, source=PrivatePortraitEventSource.CELERY.value, user_id=asset.user_id, project_id=asset.project_id, asset_id=asset.id, exc=exc)
await db.flush()
async def _delete_asset_group_remote(db: AsyncSession, *, group: PrivatePortraitAssetGroup, client: ArkPrivateAssetClient | None = None) -> None:
if not group.remote_group_id:
group.remote_delete_status = PrivatePortraitRemoteDeleteStatus.SKIPPED.value
group.remote_delete_error = None
await db.flush()
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SKIPPED.value, source=PrivatePortraitEventSource.CELERY.value, user_id=group.user_id, project_id=group.project_id, group_id=group.id, message="远程删除跳过:素材组没有 remote_group_id")
return
client = client or ArkPrivateAssetClient(for_celery=True)
now = datetime.now(timezone.utc)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CELERY.value, user_id=group.user_id, project_id=group.project_id, group_id=group.id, detail={"remote_group_id": group.remote_group_id, "remote_project_name": group.remote_project_name, "library_type": group.library_type})
try:
await client.delete_asset_group(project_name=group.remote_project_name, group_id=group.remote_group_id)
group.status = PrivatePortraitAssetGroupStatus.REMOTE_DELETED.value
group.remote_delete_status = PrivatePortraitRemoteDeleteStatus.SUCCESS.value
group.remote_deleted_at = now
group.remote_delete_error = None
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.CELERY.value, user_id=group.user_id, project_id=group.project_id, group_id=group.id, detail={"remote_group_id": group.remote_group_id, "remote_project_name": group.remote_project_name, "library_type": group.library_type})
except Exception as exc:
group.status = PrivatePortraitAssetGroupStatus.DELETE_FAILED.value
group.remote_delete_status = PrivatePortraitRemoteDeleteStatus.FAILED.value
group.remote_delete_error = str(exc)
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_FAILED.value, source=PrivatePortraitEventSource.CELERY.value, user_id=group.user_id, project_id=group.project_id, group_id=group.id, exc=exc)
await db.flush()
async def delete_project_remote(db: AsyncSession, *, project_id: str) -> None:
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CELERY.value, project_id=project_id, message="开始远程删除私域人像素材项目资源")
rows = await db.execute(select(PrivatePortraitAsset).where(PrivatePortraitAsset.project_id == project_id))
for asset in rows.scalars().all():
await delete_asset_remote(db, asset_id=asset.id)
groups = await db.execute(select(PrivatePortraitAssetGroup).where(PrivatePortraitAssetGroup.project_id == project_id))
client = ArkPrivateAssetClient(for_celery=True)
for group in groups.scalars().all():
await _delete_asset_group_remote(db, group=group, client=client)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.PROJECT_DELETE_REMOTE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.CELERY.value, project_id=project_id, message="远程删除私域人像素材项目资源完成")
await db.flush()
async def poll_due_assets_once(db: AsyncSession, *, limit: int) -> int:
now = datetime.now(timezone.utc)
rows = await db.execute(
select(PrivatePortraitAsset.id)
.where(
PrivatePortraitAsset.deleted_at.is_(None),
PrivatePortraitAsset.status == PrivatePortraitAssetStatus.PROCESSING.value,
PrivatePortraitAsset.next_poll_at.is_not(None),
PrivatePortraitAsset.next_poll_at <= now,
)
.order_by(PrivatePortraitAsset.next_poll_at.asc())
.limit(limit)
)
ids = [row[0] for row in rows.all()]
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.SYNC_DUE_ASSETS_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CELERY.value, detail={"limit": limit, "matched_count": len(ids)})
success_count = 0
failed_count = 0
for asset_id in ids:
try:
await sync_asset_status(db, user_id=None, asset_id=asset_id)
success_count += 1
except Exception as exc:
failed_count += 1
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.ASSET_POLL_FAILED.value, source=PrivatePortraitEventSource.CELERY.value, asset_id=asset_id, exc=exc)
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.SYNC_DUE_ASSETS_DONE.value, event_status=PrivatePortraitEventStatus.SUCCESS.value if failed_count == 0 else PrivatePortraitEventStatus.WARNING.value, source=PrivatePortraitEventSource.CELERY.value, detail={"matched_count": len(ids), "success_count": success_count, "failed_count": failed_count})
return len(ids)
async def recover_remote_deletes_once(db: AsyncSession, *, limit: int) -> dict[str, int]:
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.REMOTE_DELETE_RECOVERY_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CELERY.value, detail={"limit": limit})
statuses = [PrivatePortraitRemoteDeleteStatus.PENDING.value, PrivatePortraitRemoteDeleteStatus.FAILED.value]
asset_rows = await db.execute(select(PrivatePortraitAsset.id).where(PrivatePortraitAsset.remote_delete_status.in_(statuses)).order_by(PrivatePortraitAsset.updated_at.asc()).limit(limit))
asset_ids = [row[0] for row in asset_rows.all()]
for asset_id in asset_ids:
await delete_asset_remote(db, asset_id=asset_id)
remaining = max(0, limit - len(asset_ids))
group_count = 0
if remaining > 0:
group_rows = await db.execute(select(PrivatePortraitAssetGroup).where(PrivatePortraitAssetGroup.remote_delete_status.in_(statuses)).order_by(PrivatePortraitAssetGroup.updated_at.asc()).limit(remaining))
client = ArkPrivateAssetClient(for_celery=True)
groups = list(group_rows.scalars().all())
group_count = len(groups)
for group in groups:
await _delete_asset_group_remote(db, group=group, client=client)
result = {"asset_count": len(asset_ids), "group_count": group_count, "total_count": len(asset_ids) + group_count}
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.REMOTE_DELETE_RECOVERY_DONE.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.CELERY.value, detail=result)
return result