This commit is contained in:
2026-07-07 10:55:45 +08:00
23 changed files with 1552 additions and 80 deletions
@@ -54,7 +54,7 @@ from app.services.generation_history_meta_service import (
build_empty_history_meta,
)
from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls, resolve_private_portrait_references
from app.utils.id_gen import generate_id
IMAGE_DEFAULT_SIZE = "2K"
@@ -88,6 +88,22 @@ def _parse_json(text: str | None):
return None
async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[ChatGenerationTask], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
return await batch_resolve_private_portrait_reference_display_urls(
db,
{task.id: _parse_json(task.media_references) for task in tasks},
user_id=user_id,
)
async def _resolve_generation_record_reference_display_map(db: AsyncSession, records: list[GenerationRecord], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
return await batch_resolve_private_portrait_reference_display_urls(
db,
{record.id: _parse_json(record.media_references) for record in records},
user_id=user_id,
)
async def _get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
query = select(ImageEngine).where(ImageEngine.is_active == True)
if engine_id:
@@ -462,8 +478,9 @@ def record_to_out(
generated_resource_id: str | None = None,
file_name: str | None = None,
history_meta: GenerationHistoryMeta | None = None,
media_references: list[dict] | None = None,
) -> GenerationAITaskOut:
refs = _parse_json(task.media_references)
refs = media_references if media_references is not None else _parse_json(task.media_references)
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
source = GenerationHistorySourceEnum.CHAT_TASK
@@ -688,8 +705,9 @@ def generation_record_to_history_out(
project_name: str | None = None,
generated_resource_id: str | None = None,
file_name: str | None = None,
media_references: list[dict] | None = None,
) -> GenerationAIRecordHistoryItemOut:
refs = _parse_json(record.media_references)
refs = media_references if media_references is not None else _parse_json(record.media_references)
return GenerationAIRecordHistoryItemOut(
id=record.id,
source_type="generation_record",
@@ -814,6 +832,8 @@ async def list_generation_record_history_grouped_days(
source_ids=all_record_ids,
resource_type=gen_type,
)
all_records = [record for _generated_day, _day_total, rows in raw_groups for record, _project_name in rows]
reference_display_map = await _resolve_generation_record_reference_display_map(db, all_records, user_id=user_id)
groups = [
{
@@ -825,6 +845,7 @@ async def list_generation_record_history_grouped_days(
project_name,
generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"),
file_name=resource_info_map.get(record.id, {}).get("file_name"),
media_references=reference_display_map.get(record.id),
)
for record, project_name in rows
],
@@ -887,6 +908,7 @@ async def list_generation_record_history_day_items(
source_ids=[record.id for record, _project_name in rows],
resource_type=gen_type,
)
reference_display_map = await _resolve_generation_record_reference_display_map(db, [record for record, _project_name in rows], user_id=user_id)
return {
"generated_date": target_day.strftime("%Y-%m-%d"),
@@ -899,6 +921,7 @@ async def list_generation_record_history_day_items(
project_name,
generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"),
file_name=resource_info_map.get(record.id, {}).get("file_name"),
media_references=reference_display_map.get(record.id),
)
for record, project_name in rows
],
@@ -989,6 +1012,8 @@ async def list_generation_history_grouped_days(
source=source,
chat_task_ids=all_task_ids,
)
all_tasks = [task for _generated_day, _day_total, tasks in raw_groups for task in tasks]
reference_display_map = await _resolve_task_reference_display_map(db, all_tasks, user_id=user_id)
groups = [
{
@@ -1000,6 +1025,7 @@ async def list_generation_history_grouped_days(
generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"),
file_name=resource_info_map.get(task.id, {}).get("file_name"),
history_meta=history_meta_map.get(task.id),
media_references=reference_display_map.get(task.id),
)
for task in tasks
],
@@ -1082,6 +1108,7 @@ async def list_generation_history_day_items(
source=source,
chat_task_ids=task_ids,
)
reference_display_map = await _resolve_task_reference_display_map(db, tasks, user_id=user_id)
return {
"generated_date": target_day.strftime("%Y-%m-%d"),
@@ -1094,6 +1121,7 @@ async def list_generation_history_day_items(
generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"),
file_name=resource_info_map.get(task.id, {}).get("file_name"),
history_meta=history_meta_map.get(task.id),
media_references=reference_display_map.get(task.id),
)
for task in tasks
],
@@ -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)
@@ -17,7 +17,7 @@ from app.enums.private_portrait import (
PrivatePortraitProjectStatus,
PrivatePortraitRemoteDeleteStatus,
)
from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitAssetGroup, PrivatePortraitProject, PrivatePortraitValidateSession
from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitAssetGroup, PrivatePortraitProject
from app.schemas.private_portrait import PrivatePortraitProjectCreate, PrivatePortraitProjectOut, PrivatePortraitProjectUpdate
from app.services.operation_log_service import log_operation_event
from app.utils.id_gen import generate_id
@@ -74,19 +74,19 @@ async def create_project(db: AsyncSession, *, user_id: str, payload: PrivatePort
name_slug=slug,
remote_project_name=PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME,
description=payload.description,
status=PrivatePortraitProjectStatus.ACTIVE.value,
status=PrivatePortraitProjectStatus.VALIDATING.value,
)
db.add(project)
await db.flush()
log_operation_event(
domain=DOMAIN,
event_type=PrivatePortraitEventType.PROJECT_CREATE.value,
event_status=PrivatePortraitEventStatus.SUCCESS.value,
event_status=PrivatePortraitEventStatus.PENDING.value,
source=PrivatePortraitEventSource.API.value,
user_id=user_id,
project_id=project.id,
message="创建真人素材项目",
detail={"name": project.name, "remote_project_name": project.remote_project_name},
message="创建待认证真人素材项目",
detail={"name": project.name, "remote_project_name": project.remote_project_name, "status": project.status},
)
return project
@@ -109,7 +109,12 @@ async def update_project(db: AsyncSession, *, user_id: str, project_id: str, pay
if payload.description is not None:
project.description = payload.description
if payload.status is not None:
if payload.status not in {PrivatePortraitProjectStatus.ACTIVE.value}:
allowed_statuses = {
PrivatePortraitProjectStatus.VALIDATING.value,
PrivatePortraitProjectStatus.ACTIVE.value,
PrivatePortraitProjectStatus.VALIDATE_FAILED.value,
}
if payload.status not in allowed_statuses:
raise HTTPException(status_code=400, detail="项目状态不支持")
project.status = payload.status
await db.flush()
@@ -133,7 +138,15 @@ async def update_project(db: AsyncSession, *, user_id: str, project_id: str, pay
return project
async def list_projects(db: AsyncSession, *, user_id: str | None, page: int = 1, page_size: int = 20, keyword: str | None = None, status: str | None = None) -> tuple[list[PrivatePortraitProject], int]:
async def list_projects(
db: AsyncSession,
*,
user_id: str | None,
page: int = 1,
page_size: int = 20,
keyword: str | None = None,
status: str | None = None,
) -> tuple[list[PrivatePortraitProject], int]:
page = max(1, page)
page_size = min(max(1, page_size), 100)
filters = [PrivatePortraitProject.deleted_at.is_(None)]
@@ -4,7 +4,7 @@ from copy import deepcopy
from typing import Any
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy import or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.private_portrait import (
@@ -58,6 +58,162 @@ def _normalize_ref_type(value: Any) -> str:
return str(value or "").strip().lower()
def _remote_asset_id_from_asset_uri(url: Any) -> str | None:
value = str(url or "").strip()
if not value.startswith(PRIVATE_PORTRAIT_ASSET_URI_PREFIX):
return None
remote_asset_id = value[len(PRIVATE_PORTRAIT_ASSET_URI_PREFIX):].strip()
return remote_asset_id or None
def _asset_display_url(asset: PrivatePortraitAsset) -> str | None:
# preview_url 是本地上传预览,remote_url 是火山 GetAsset 返回的远程资源 URLsource_url 是兜底公网上传地址。
return asset.preview_url or asset.remote_url or asset.source_url or None
def _fill_private_portrait_reference_display_fields(ref: Any, asset: PrivatePortraitAsset) -> None:
provider_url = str(_ref_get(ref, "provider_url") or _ref_get(ref, "url") or "").strip()
if provider_url.startswith(PRIVATE_PORTRAIT_ASSET_URI_PREFIX):
_ref_set(ref, "provider_url", provider_url)
display_url = _asset_display_url(asset)
if display_url:
# 返回给前端的 url 必须可预览;供应商专用 asset:// 保留到 provider_url,避免管理后台和客户端展示黑图。
_ref_set(ref, "url", display_url)
_ref_set(ref, "display_url", display_url)
_ref_set(ref, "preview_url", display_url)
_ref_set(ref, "source", PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value)
_ref_set(ref, "private_asset_id", asset.id)
if asset.remote_asset_id:
_ref_set(ref, "remote_asset_id", asset.remote_asset_id)
expected_ref_type = _ASSET_TYPE_TO_REFERENCE_TYPE.get(asset.asset_type)
if expected_ref_type:
_ref_set(ref, "type", expected_ref_type)
if not _ref_get(ref, "name") and asset.name:
_ref_set(ref, "name", asset.name)
async def resolve_private_portrait_reference_display_urls(
db: AsyncSession,
media_references: list[Any] | None,
*,
user_id: str | None = None,
) -> list[Any] | None:
"""把历史响应里的 asset:// 引用补成前端可预览 URL。
生成任务入库时 url 使用 asset://remote_asset_id 传给供应商;但客户端/管理后台展示不能直接用
asset://。这里批量根据 private_asset_id 或 asset://remote_asset_id 查本地素材,并把响应中的 url
改成 preview_url/remote_url/source_url,同时保留 provider_url=asset://... 供排查。
"""
if not media_references:
return media_references
refs = deepcopy(media_references)
private_asset_ids: list[str] = []
remote_asset_ids: list[str] = []
for ref in refs:
source = _ref_get(ref, "source")
private_asset_id = _ref_get(ref, "private_asset_id")
remote_asset_id = _ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url"))
if source == PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value or remote_asset_id:
if private_asset_id:
private_asset_ids.append(str(private_asset_id))
if remote_asset_id:
remote_asset_ids.append(str(remote_asset_id))
private_asset_ids = list(dict.fromkeys(private_asset_ids))
remote_asset_ids = list(dict.fromkeys(remote_asset_ids))
if not private_asset_ids and not remote_asset_ids:
return refs
filters = []
if private_asset_ids:
filters.append(PrivatePortraitAsset.id.in_(private_asset_ids))
if remote_asset_ids:
filters.append(PrivatePortraitAsset.remote_asset_id.in_(remote_asset_ids))
stmt = select(PrivatePortraitAsset).where(or_(*filters))
if user_id is not None:
stmt = stmt.where(PrivatePortraitAsset.user_id == user_id)
rows = await db.execute(stmt)
assets = list(rows.scalars().all())
by_id = {asset.id: asset for asset in assets}
by_remote_id = {asset.remote_asset_id: asset for asset in assets if asset.remote_asset_id}
for ref in refs:
private_asset_id = str(_ref_get(ref, "private_asset_id") or "").strip()
remote_asset_id = str(_ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url")) or "").strip()
asset = by_id.get(private_asset_id) or by_remote_id.get(remote_asset_id)
if not asset:
continue
_fill_private_portrait_reference_display_fields(ref, asset)
return refs
async def batch_resolve_private_portrait_reference_display_urls(
db: AsyncSession,
references_by_key: dict[Any, list[Any] | None],
*,
user_id: str | None = None,
) -> dict[Any, list[Any] | None]:
if not references_by_key:
return {}
copied: dict[Any, list[Any] | None] = {
key: deepcopy(refs) if refs else refs
for key, refs in references_by_key.items()
}
private_asset_ids: list[str] = []
remote_asset_ids: list[str] = []
for refs in copied.values():
if not refs:
continue
for ref in refs:
source = _ref_get(ref, "source")
private_asset_id = _ref_get(ref, "private_asset_id")
remote_asset_id = _ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url"))
if source == PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value or remote_asset_id:
if private_asset_id:
private_asset_ids.append(str(private_asset_id))
if remote_asset_id:
remote_asset_ids.append(str(remote_asset_id))
private_asset_ids = list(dict.fromkeys(private_asset_ids))
remote_asset_ids = list(dict.fromkeys(remote_asset_ids))
if not private_asset_ids and not remote_asset_ids:
return copied
filters = []
if private_asset_ids:
filters.append(PrivatePortraitAsset.id.in_(private_asset_ids))
if remote_asset_ids:
filters.append(PrivatePortraitAsset.remote_asset_id.in_(remote_asset_ids))
stmt = select(PrivatePortraitAsset).where(or_(*filters))
if user_id is not None:
stmt = stmt.where(PrivatePortraitAsset.user_id == user_id)
rows = await db.execute(stmt)
assets = list(rows.scalars().all())
by_id = {asset.id: asset for asset in assets}
by_remote_id = {asset.remote_asset_id: asset for asset in assets if asset.remote_asset_id}
for refs in copied.values():
if not refs:
continue
for ref in refs:
private_asset_id = str(_ref_get(ref, "private_asset_id") or "").strip()
remote_asset_id = str(_ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url")) or "").strip()
asset = by_id.get(private_asset_id) or by_remote_id.get(remote_asset_id)
if not asset:
continue
_fill_private_portrait_reference_display_fields(ref, asset)
return copied
async def resolve_private_portrait_references(
db: AsyncSession,
*,
@@ -123,10 +279,16 @@ async def resolve_private_portrait_references(
if expected_ref_type and ref_type and ref_type != expected_ref_type:
raise HTTPException(status_code=400, detail=f"真人素材类型不匹配:引用为 {ref_type},素材为 {expected_ref_type}")
provider_url = f"{PRIVATE_PORTRAIT_ASSET_URI_PREFIX}{asset.remote_asset_id}"
_ref_set(ref, "source", PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value)
_ref_set(ref, "private_asset_id", asset.id)
_ref_set(ref, "remote_asset_id", asset.remote_asset_id)
_ref_set(ref, "url", f"{PRIVATE_PORTRAIT_ASSET_URI_PREFIX}{asset.remote_asset_id}")
_ref_set(ref, "url", provider_url)
_ref_set(ref, "provider_url", provider_url)
display_url = _asset_display_url(asset)
if display_url:
_ref_set(ref, "display_url", display_url)
_ref_set(ref, "preview_url", display_url)
if expected_ref_type:
_ref_set(ref, "type", expected_ref_type)
if not _ref_get(ref, "name") and asset.name: