from __future__ import annotations from copy import deepcopy from typing import Any from fastapi import HTTPException from sqlalchemy import or_, select from sqlalchemy.ext.asyncio import AsyncSession from app.enums.private_portrait import ( PRIVATE_PORTRAIT_ASSET_URI_PREFIX, PrivatePortraitAssetStatus, PRIVATE_PORTRAIT_ENABLED_ASSET_TYPES, PrivatePortraitAssetType, PrivatePortraitEventSource, PrivatePortraitEventStatus, PrivatePortraitEventType, PrivatePortraitReferenceSource, ) from app.models.private_portrait import PrivatePortraitAsset from app.services.operation_log_service import log_operation_error, log_operation_event DOMAIN = "private_portrait" _ASSET_TYPE_TO_REFERENCE_TYPE = { PrivatePortraitAssetType.IMAGE.value: "image", PrivatePortraitAssetType.VIDEO.value: "video", } _SUPPORTED_GEN_TYPES = {"image", "video"} def _ref_get(ref: Any, key: str) -> Any: if isinstance(ref, dict): return ref.get(key) return getattr(ref, key, None) def _ref_set(ref: Any, key: str, value: Any) -> None: if isinstance(ref, dict): ref[key] = value else: setattr(ref, key, value) def _normalize_gen_type(gen_type: str | None) -> str | None: value = (gen_type or "").strip().lower() if not value: return None if value not in _SUPPORTED_GEN_TYPES: raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video") return value 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 返回的远程资源 URL;source_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 asset.asset_type == PrivatePortraitAssetType.VIDEO.value and asset.video_duration is not None: _ref_set(ref, "duration", asset.video_duration) 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, *, user_id: str, media_references: list[Any] | None, gen_type: str | None = None, ) -> list[Any] | None: """Resolve private portrait references before dispatching a generation task. generation_ai_service.py 和 generation_task_factory_service.py 都会传 gen_type。 这里保留该参数用于兼容调用方,并做基础校验,避免接口因函数签名不一致直接 500。 """ normalized_gen_type = _normalize_gen_type(gen_type) if not media_references: return media_references refs = deepcopy(media_references) ids = [ str(_ref_get(ref, "private_asset_id")) for ref in refs if _ref_get(ref, "source") == PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value and _ref_get(ref, "private_asset_id") ] ids = list(dict.fromkeys(ids)) if not ids: return refs log_operation_event( domain=DOMAIN, event_type=PrivatePortraitEventType.REFERENCE_RESOLVE_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.SERVICE.value, user_id=user_id, detail={"private_asset_ids": ids, "gen_type": normalized_gen_type}, ) try: rows = await db.execute(select(PrivatePortraitAsset).where(PrivatePortraitAsset.id.in_(ids))) asset_map = {asset.id: asset for asset in rows.scalars().all()} for ref in refs: if _ref_get(ref, "source") != PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value: continue asset_id = str(_ref_get(ref, "private_asset_id") or "") if not asset_id: raise HTTPException(status_code=400, detail="私域人像素材引用缺少 private_asset_id") asset = asset_map.get(asset_id) if not asset: raise HTTPException(status_code=400, detail="私域人像素材不存在") if asset.user_id != user_id: raise HTTPException(status_code=403, detail="私域人像素材不属于当前用户") if asset.deleted_at is not None: raise HTTPException(status_code=400, detail="私域人像素材已删除") if asset.asset_type not in PRIVATE_PORTRAIT_ENABLED_ASSET_TYPES: raise HTTPException(status_code=400, detail="Audio 暂未开放,当前仅支持 Image/Video 私域素材") if asset.status != PrivatePortraitAssetStatus.ACTIVE.value: raise HTTPException(status_code=400, detail=f"私域人像素材状态为 {asset.status},Active 后才可用于生成") if not asset.remote_asset_id: raise HTTPException(status_code=400, detail="私域人像素材缺少远程 AssetId") expected_ref_type = _ASSET_TYPE_TO_REFERENCE_TYPE.get(asset.asset_type) ref_type = _normalize_ref_type(_ref_get(ref, "type")) 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", 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 asset.asset_type == PrivatePortraitAssetType.VIDEO.value: if asset.video_duration is None: raise HTTPException(status_code=400, detail="私域视频素材缺少 video_duration,不能用于生成") _ref_set(ref, "duration", asset.video_duration) if not _ref_get(ref, "name") and asset.name: _ref_set(ref, "name", asset.name) log_operation_event( domain=DOMAIN, event_type=PrivatePortraitEventType.REFERENCE_RESOLVE_SUCCESS.value, event_status=PrivatePortraitEventStatus.SUCCESS.value, source=PrivatePortraitEventSource.SERVICE.value, user_id=user_id, detail={"count": len(ids), "gen_type": normalized_gen_type}, ) return refs except Exception as exc: log_operation_error( domain=DOMAIN, event_type=PrivatePortraitEventType.REFERENCE_RESOLVE_FAILED.value, source=PrivatePortraitEventSource.SERVICE.value, user_id=user_id, detail={"private_asset_ids": ids, "gen_type": normalized_gen_type}, exc=exc, ) raise