真人素材库修复相关BUG

This commit is contained in:
2026-07-07 10:48:52 +08:00
parent ab995ce3cf
commit 126754c933
20 changed files with 630 additions and 79 deletions
@@ -23,6 +23,7 @@ from app.enums.private_portrait import (
PrivatePortraitEventSource,
PrivatePortraitEventStatus,
PrivatePortraitEventType,
PrivatePortraitProjectStatus,
PrivatePortraitRemoteDeleteStatus,
PrivatePortraitValidateSessionStatus,
)
@@ -165,8 +166,69 @@ def asset_to_out(asset: PrivatePortraitAsset, *, project_name: str | None = None
)
async def _get_existing_active_group(db: AsyncSession, *, project_id: str) -> PrivatePortraitAssetGroup | None:
return (
await db.execute(
select(PrivatePortraitAssetGroup)
.where(
PrivatePortraitAssetGroup.project_id == project_id,
PrivatePortraitAssetGroup.status == PrivatePortraitAssetGroupStatus.ACTIVE.value,
PrivatePortraitAssetGroup.deleted_at.is_(None),
)
.order_by(PrivatePortraitAssetGroup.created_at.desc())
.limit(1)
)
).scalar_one_or_none()
async def _ensure_project_can_validate(db: AsyncSession, *, project: PrivatePortraitProject) -> PrivatePortraitValidateSession | None:
active_group = await _get_existing_active_group(db, project_id=project.id)
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)
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,
@@ -192,11 +254,11 @@ async def create_validate_session(db: AsyncSession, *, user_id: str, project_id:
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:
@@ -209,6 +271,9 @@ async def get_validate_session(db: AsyncSession, *, user_id: str | None, 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
@@ -218,9 +283,13 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
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
@@ -229,10 +298,22 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
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)
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")
@@ -242,7 +323,8 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
session.status = PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value
session.raw_response_json = _json(resp)
project = (await db.execute(select(PrivatePortraitProject).where(PrivatePortraitProject.id == session.project_id).limit(1))).scalar_one()
if not project:
raise RuntimeError("真人素材项目不存在")
remote_group_name = _remote_group_name(session.user_id, project.name)
group = PrivatePortraitAssetGroup(
id=generate_id(),
@@ -256,6 +338,7 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
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)
@@ -269,12 +352,16 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
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) -> PrivatePortraitAssetGroup:
project = await get_user_project(db, user_id=user_id, project_id=project_id)
if project.status != PrivatePortraitProjectStatus.ACTIVE.value:
raise HTTPException(status_code=400, detail="请先完成真人授权认证,再上传素材")
result = await db.execute(select(PrivatePortraitAssetGroup).where(PrivatePortraitAssetGroup.user_id == user_id, PrivatePortraitAssetGroup.project_id == project_id, 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:
@@ -284,6 +371,8 @@ async def get_project_active_group(db: AsyncSession, *, user_id: str, project_id
async def create_asset(db: AsyncSession, *, user_id: str, project_id: str, payload: PrivatePortraitAssetCreate) -> PrivatePortraitAsset:
project = await get_user_project(db, user_id=user_id, project_id=project_id)
if project.status != PrivatePortraitProjectStatus.ACTIVE.value:
raise HTTPException(status_code=400, detail="项目正在真人认证或认证未通过,不能上传素材")
if payload.asset_type != PrivatePortraitAssetType.IMAGE.value:
raise HTTPException(status_code=400, detail="第一版真人素材库仅开放 Image 图片素材")
user = await _lock_user_for_upload(db, user_id=user_id)