真人素材库授权BUG修复 3

This commit is contained in:
2026-07-06 16:46:00 +08:00
parent 732cdd6ed1
commit df251b83f5
6 changed files with 106 additions and 73 deletions
-1
View File
@@ -3,7 +3,6 @@ APP_NAME=VideoGen API
APP_VERSION=1.0.0 APP_VERSION=1.0.0
DEBUG=false DEBUG=false
SECRET_KEY=local-dev-secret-key-not-for-production SECRET_KEY=local-dev-secret-key-not-for-production
APP_ENV=test
# Database (PostgreSQL) # Database (PostgreSQL)
#DATABASE_URL=postgresql+asyncpg://videogen_test:Yr7kM7kDj75izCiA@180.184.42.66:5432/videogen_test #DATABASE_URL=postgresql+asyncpg://videogen_test:Yr7kM7kDj75izCiA@180.184.42.66:5432/videogen_test
-2
View File
@@ -7,8 +7,6 @@ class Settings(BaseSettings):
APP_NAME: str = "VideoGen API" APP_NAME: str = "VideoGen API"
APP_VERSION: str = "1.0.0" APP_VERSION: str = "1.0.0"
DEBUG: bool = False DEBUG: bool = False
# 运行环境,用于生成火山私域真人素材 ProjectNamelocal/test/online。
APP_ENV: str = "local"
SECRET_KEY: str = "change-me" SECRET_KEY: str = "change-me"
DATABASE_URL: str = "sqlite+aiosqlite:///./videogen.db" DATABASE_URL: str = "sqlite+aiosqlite:///./videogen.db"
+4 -11
View File
@@ -5,15 +5,9 @@ from enum import Enum
# 用户真人素材图片默认上限。users.private_portrait_image_limit = 0 表示关闭模块;>0 表示启用并限制总量。 # 用户真人素材图片默认上限。users.private_portrait_image_limit = 0 表示关闭模块;>0 表示启用并限制总量。
PRIVATE_PORTRAIT_DEFAULT_IMAGE_LIMIT = 5 PRIVATE_PORTRAIT_DEFAULT_IMAGE_LIMIT = 5
# ProjectName 由服务层按 {env}-{user_id}-{项目名slug} 生成并快照到 private_portrait_projects.remote_project_name。 # 火山 Ark 私域真人素材 ProjectName:火山侧项目空间固定使用 default,并快照到各业务表 remote_project_name。
PRIVATE_PORTRAIT_PROJECT_ENV_LOCAL = "local" # 用户/项目隔离依赖本地 project_id 和火山返回的 Asset Group ID,不再动态拼接 ProjectName。
PRIVATE_PORTRAIT_PROJECT_ENV_TEST = "test" PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME = "default"
PRIVATE_PORTRAIT_PROJECT_ENV_ONLINE = "online"
PRIVATE_PORTRAIT_PROJECT_ENV_VALUES = {
PRIVATE_PORTRAIT_PROJECT_ENV_LOCAL,
PRIVATE_PORTRAIT_PROJECT_ENV_TEST,
PRIVATE_PORTRAIT_PROJECT_ENV_ONLINE,
}
PRIVATE_PORTRAIT_GROUP_TYPE = "LivenessFace" PRIVATE_PORTRAIT_GROUP_TYPE = "LivenessFace"
PRIVATE_PORTRAIT_VERIFY_TYPE = "real_time" PRIVATE_PORTRAIT_VERIFY_TYPE = "real_time"
@@ -132,9 +126,8 @@ class PrivatePortraitEventType(str, Enum):
PROJECT_CREATE = "PROJECT_CREATE" PROJECT_CREATE = "PROJECT_CREATE"
PROJECT_UPDATE = "PROJECT_UPDATE" PROJECT_UPDATE = "PROJECT_UPDATE"
PROJECT_DELETE = "PROJECT_DELETE" PROJECT_DELETE = "PROJECT_DELETE"
PROJECT_REMOTE_NAME_LOCKED = "PROJECT_REMOTE_NAME_LOCKED"
VALIDATE_SESSION_CREATE = "VALIDATE_SESSION_CREATE" VALIDATE_SESSION_CREATE = "VALIDATE_SESSION_CREATE"
VALIDATE_SESSION_CREATE_FAILED = "VALIDATE_SESSION_CREATE_FAILED"
VALIDATE_CALLBACK_RECEIVED = "VALIDATE_CALLBACK_RECEIVED" VALIDATE_CALLBACK_RECEIVED = "VALIDATE_CALLBACK_RECEIVED"
VALIDATE_CALLBACK_SUCCESS = "VALIDATE_CALLBACK_SUCCESS" VALIDATE_CALLBACK_SUCCESS = "VALIDATE_CALLBACK_SUCCESS"
VALIDATE_CALLBACK_FAILED = "VALIDATE_CALLBACK_FAILED" VALIDATE_CALLBACK_FAILED = "VALIDATE_CALLBACK_FAILED"
@@ -27,6 +27,29 @@ class ArkPrivateAssetClientError(RuntimeError):
pass pass
class ArkPrivateAssetRemoteError(ArkPrivateAssetClientError):
def __init__(self, *, action: str, code: str, message: str, request_id: str | None = None, raw: dict[str, Any] | None = None):
self.action = action
self.code = code
self.message = message
self.request_id = request_id
self.raw = raw or {}
super().__init__(f"{action} 调用失败:{code} {message}" + (f" RequestId={request_id}" if request_id else ""))
def _remote_error_http_status(code: str) -> int:
normalized = (code or "").lower()
if "notfound" in normalized or normalized.startswith("not_found"):
return 404
if "invalid" in normalized or "missing" in normalized or "unsupported" in normalized or "limit" in normalized:
return 400
if "unauthorized" in normalized or "forbidden" in normalized or "permission" in normalized or "denied" in normalized:
return 403
if "throttl" in normalized or "rate" in normalized:
return 429
return 502
class ArkPrivateAssetClient: class ArkPrivateAssetClient:
"""火山 Ark 私域真人人像素材 API Client。只做 AK/SK 鉴权调用与响应标准化。""" """火山 Ark 私域真人人像素材 API Client。只做 AK/SK 鉴权调用与响应标准化。"""
@@ -109,6 +132,24 @@ class ArkPrivateAssetClient:
remote_request_id=result.get("RequestId") or result.get("request_id"), remote_request_id=result.get("RequestId") or result.get("request_id"),
) )
return result return result
except ArkPrivateAssetRemoteError as exc:
log_remote_api_event(
domain=DOMAIN,
remote_action=action_value,
event_type=PrivatePortraitEventType.ARK_API_CALL_FAILED.value,
event_status=PrivatePortraitEventStatus.FAILED.value,
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
request=payload,
response=exc.raw,
remote_request_id=exc.request_id,
remote_code=exc.code,
remote_message=exc.message,
)
if self.for_celery:
raise
raise HTTPException(status_code=_remote_error_http_status(exc.code), detail={"message": exc.message, "code": exc.code, "request_id": exc.request_id}) from exc
except HTTPException:
raise
except Exception as exc: except Exception as exc:
log_remote_api_event( log_remote_api_event(
domain=DOMAIN, domain=DOMAIN,
@@ -119,7 +160,9 @@ class ArkPrivateAssetClient:
request=payload, request=payload,
remote_message=str(exc), remote_message=str(exc),
) )
raise if self.for_celery:
raise
raise HTTPException(status_code=502, detail=f"火山私域素材接口调用失败:{self._exception_message(exc)}") from exc
def _call_sync(self, action: ArkPrivatePortraitAction, payload: dict[str, Any]) -> dict[str, Any]: def _call_sync(self, action: ArkPrivatePortraitAction, payload: dict[str, Any]) -> dict[str, Any]:
try: try:
@@ -140,6 +183,7 @@ class ArkPrivateAssetClient:
"https", "https",
) )
# volcengine SDK 的 ApiInfo.query 必须是 dict,不能传 "Action=xxx&Version=xxx" 字符串。 # volcengine SDK 的 ApiInfo.query 必须是 dict,不能传 "Action=xxx&Version=xxx" 字符串。
# volcengine SDK 签名时要求 body 是 bytes。
# ServiceInfo 默认 scheme='http',这里必须显式传 https,避免请求走 http://host:80。 # ServiceInfo 默认 scheme='http',这里必须显式传 https,避免请求走 http://host:80。
api_info = { api_info = {
action.value: ApiInfo( action.value: ApiInfo(
@@ -151,26 +195,22 @@ class ArkPrivateAssetClient:
) )
} }
service = Service(service_info, api_info) service = Service(service_info, api_info)
# volcengine SDK 在签名阶段会对 request.body 做 hashlib.sha256(body)。
# Python 3 下 hashlib.sha256 只能接收 bytes/bytearray,不能接收 dict 或 str。
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8") body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
try: try:
raw = service.json(action.value, {}, body) raw = service.json(action.value, {}, body)
except Exception as exc: except Exception as exc:
resp = self._extract_response_from_exception(exc)
if resp:
self._raise_remote_error_if_any(action.value, resp)
message = self._exception_message(exc) message = self._exception_message(exc)
if "ConnectTimeout" in message or "timed out" in message or "Connection" in message: if "ConnectTimeout" in message or "timed out" in message or "Connection" in message:
raise HTTPException(status_code=502, detail=f"火山私域素材接口连接失败:{message}") from exc raise ArkPrivateAssetClientError(f"{action.value} 连接火山私域素材接口失败:{message}") from exc
raise ArkPrivateAssetClientError(f"{action.value} 请求火山私域素材接口异常:{message}") from exc raise ArkPrivateAssetClientError(f"{action.value} 请求火山私域素材接口异常:{message}") from exc
resp = self._normalize_response(raw) resp = self._normalize_response(raw)
self._raise_remote_error_if_any(action.value, resp)
metadata = resp.get("ResponseMetadata") if isinstance(resp, dict) else None metadata = resp.get("ResponseMetadata") if isinstance(resp, dict) else None
error = metadata.get("Error") if isinstance(metadata, dict) else None
request_id = metadata.get("RequestId") if isinstance(metadata, dict) else None request_id = metadata.get("RequestId") if isinstance(metadata, dict) else None
if error:
code = error.get("Code") or "ArkPrivateAssetError"
message = error.get("Message") or str(error)
raise ArkPrivateAssetClientError(f"{action.value} 调用失败:{code} {message}")
if isinstance(resp, dict) and isinstance(resp.get("Result"), dict): if isinstance(resp, dict) and isinstance(resp.get("Result"), dict):
result = dict(resp["Result"]) result = dict(resp["Result"])
if request_id: if request_id:
@@ -182,6 +222,36 @@ class ArkPrivateAssetClient:
return resp return resp
return {"raw": resp, "RequestId": request_id} return {"raw": resp, "RequestId": request_id}
@classmethod
def _raise_remote_error_if_any(cls, action: str, resp: dict[str, Any]) -> None:
metadata = resp.get("ResponseMetadata") if isinstance(resp, dict) else None
error = metadata.get("Error") if isinstance(metadata, dict) else None
if not error:
return
request_id = metadata.get("RequestId") if isinstance(metadata, dict) else None
code = error.get("Code") or "ArkPrivateAssetError"
message = error.get("Message") or str(error)
raise ArkPrivateAssetRemoteError(action=action, code=code, message=message, request_id=request_id, raw=resp)
@classmethod
def _extract_response_from_exception(cls, exc: Exception) -> dict[str, Any] | None:
# volcengine SDK 在 HTTP 非 2xx 时会 raise Exception(resp.text.encode("utf-8")),这里把 bytes JSON 还原,避免业务错误变 500。
if not exc.args:
return None
raw = exc.args[0]
if isinstance(raw, (bytes, bytearray)):
raw = raw.decode("utf-8", errors="ignore")
if not isinstance(raw, str):
return None
raw = raw.strip()
if not raw or not raw.startswith("{"):
return None
try:
obj = json.loads(raw)
except json.JSONDecodeError:
return None
return obj if isinstance(obj, dict) else None
@staticmethod @staticmethod
def _exception_message(exc: Exception) -> str: def _exception_message(exc: Exception) -> str:
if exc.args: if exc.args:
@@ -52,6 +52,17 @@ def _loads(data: str | None) -> Any:
return None 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: def _public_url(url: str) -> str:
if url.startswith(("http://", "https://")): if url.startswith(("http://", "https://")):
return url return url
@@ -177,9 +188,9 @@ async def create_validate_session(db: AsyncSession, *, user_id: str, project_id:
return session return session
except Exception as exc: except Exception as exc:
session.status = PrivatePortraitValidateSessionStatus.FAILED.value session.status = PrivatePortraitValidateSessionStatus.FAILED.value
session.error_message = str(exc) session.error_message = _exception_message(exc)
await db.flush() await db.flush()
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_GET_RESULT_FAILED.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, session_id=session.id, exc=exc) 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 raise
@@ -252,7 +263,7 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
return session return session
except Exception as exc: except Exception as exc:
session.status = PrivatePortraitValidateSessionStatus.FAILED.value session.status = PrivatePortraitValidateSessionStatus.FAILED.value
session.error_message = str(exc) session.error_message = _exception_message(exc)
await db.flush() 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) 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 raise
@@ -314,7 +325,7 @@ async def create_asset(db: AsyncSession, *, user_id: str, project_id: str, paylo
return asset return asset
except Exception as exc: except Exception as exc:
asset.status = PrivatePortraitAssetStatus.FAILED.value asset.status = PrivatePortraitAssetStatus.FAILED.value
asset.error_message = str(exc) asset.error_message = _exception_message(exc)
await db.flush() 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) 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 raise
@@ -7,9 +7,8 @@ from fastapi import HTTPException
from sqlalchemy import case, func, select, update from sqlalchemy import case, func, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.private_portrait import ( from app.enums.private_portrait import (
PRIVATE_PORTRAIT_PROJECT_ENV_VALUES, PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME,
PrivatePortraitAssetGroupStatus, PrivatePortraitAssetGroupStatus,
PrivatePortraitAssetStatus, PrivatePortraitAssetStatus,
PrivatePortraitEventSource, PrivatePortraitEventSource,
@@ -28,26 +27,12 @@ DOMAIN = "private_portrait"
def _safe_slug(value: str, *, max_length: int = 80) -> str: def _safe_slug(value: str, *, max_length: int = 80) -> str:
value = (value or "").strip().lower() value = (value or "").strip().lower()
# 先保留常见英文数字连字符;中文等字符统一转 _,避免火山 ProjectName 字符限制不明确导致失败 # 先保留常见英文数字连字符;中文等字符统一转 _,仅用于本地项目 slug
value = re.sub(r"[^a-z0-9_-]+", "_", value) value = re.sub(r"[^a-z0-9_-]+", "_", value)
value = re.sub(r"_+", "_", value).strip("_-") value = re.sub(r"_+", "_", value).strip("_-")
return (value[:max_length] or "project") return (value[:max_length] or "project")
def get_private_portrait_env() -> str:
env = str(getattr(settings, "APP_ENV", "local") or "local").strip().lower()
if env not in PRIVATE_PORTRAIT_PROJECT_ENV_VALUES:
env = "local"
return env
def build_remote_project_name(*, user_id: str, project_name: str) -> tuple[str, str]:
slug = _safe_slug(project_name)
user_part = _safe_slug(user_id, max_length=40)
env = get_private_portrait_env()
return f"{env}-{user_part}-{slug}"[:256], slug
def project_to_out(project: PrivatePortraitProject, *, include_user: bool = False) -> PrivatePortraitProjectOut: def project_to_out(project: PrivatePortraitProject, *, include_user: bool = False) -> PrivatePortraitProjectOut:
return PrivatePortraitProjectOut( return PrivatePortraitProjectOut(
id=project.id, id=project.id,
@@ -81,13 +66,13 @@ async def get_user_project(db: AsyncSession, *, user_id: str, project_id: str) -
async def create_project(db: AsyncSession, *, user_id: str, payload: PrivatePortraitProjectCreate) -> PrivatePortraitProject: async def create_project(db: AsyncSession, *, user_id: str, payload: PrivatePortraitProjectCreate) -> PrivatePortraitProject:
remote_project_name, slug = build_remote_project_name(user_id=user_id, project_name=payload.name) slug = _safe_slug(payload.name)
project = PrivatePortraitProject( project = PrivatePortraitProject(
id=generate_id(), id=generate_id(),
user_id=user_id, user_id=user_id,
name=payload.name.strip(), name=payload.name.strip(),
name_slug=slug, name_slug=slug,
remote_project_name=remote_project_name, remote_project_name=PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME,
description=payload.description, description=payload.description,
status=PrivatePortraitProjectStatus.ACTIVE.value, status=PrivatePortraitProjectStatus.ACTIVE.value,
) )
@@ -106,13 +91,6 @@ async def create_project(db: AsyncSession, *, user_id: str, payload: PrivatePort
return project return project
async def _project_has_remote_resources(db: AsyncSession, *, project_id: str) -> bool:
session_count = (await db.execute(select(func.count(PrivatePortraitValidateSession.id)).where(PrivatePortraitValidateSession.project_id == project_id))).scalar_one() or 0
group_count = (await db.execute(select(func.count(PrivatePortraitAssetGroup.id)).where(PrivatePortraitAssetGroup.project_id == project_id, PrivatePortraitAssetGroup.remote_group_id.is_not(None)))).scalar_one() or 0
asset_count = (await db.execute(select(func.count(PrivatePortraitAsset.id)).where(PrivatePortraitAsset.project_id == project_id, PrivatePortraitAsset.remote_asset_id.is_not(None)))).scalar_one() or 0
return bool(session_count or group_count or asset_count)
async def update_project(db: AsyncSession, *, user_id: str, project_id: str, payload: PrivatePortraitProjectUpdate) -> PrivatePortraitProject: async def update_project(db: AsyncSession, *, user_id: str, project_id: str, payload: PrivatePortraitProjectUpdate) -> PrivatePortraitProject:
project = await get_user_project(db, user_id=user_id, project_id=project_id) project = await get_user_project(db, user_id=user_id, project_id=project_id)
before = { before = {
@@ -122,28 +100,12 @@ async def update_project(db: AsyncSession, *, user_id: str, project_id: str, pay
"description": project.description, "description": project.description,
"status": project.status, "status": project.status,
} }
remote_name_locked = False
if payload.name is not None: if payload.name is not None:
new_name = payload.name.strip() new_name = payload.name.strip()
if new_name and new_name != project.name: if new_name and new_name != project.name:
has_remote = await _project_has_remote_resources(db, project_id=project.id)
project.name = new_name project.name = new_name
if not has_remote: project.name_slug = _safe_slug(new_name)
remote_project_name, slug = build_remote_project_name(user_id=user_id, project_name=new_name) project.remote_project_name = PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME
project.name_slug = slug
project.remote_project_name = remote_project_name
else:
remote_name_locked = True
log_operation_event(
domain=DOMAIN,
event_type=PrivatePortraitEventType.PROJECT_REMOTE_NAME_LOCKED.value,
event_status=PrivatePortraitEventStatus.SKIPPED.value,
source=PrivatePortraitEventSource.API.value,
user_id=user_id,
project_id=project.id,
message="项目已有远程资源,仅修改展示名,remote_project_name 保持不变",
detail={"remote_project_name": project.remote_project_name},
)
if payload.description is not None: if payload.description is not None:
project.description = payload.description project.description = payload.description
if payload.status is not None: if payload.status is not None:
@@ -166,7 +128,7 @@ async def update_project(db: AsyncSession, *, user_id: str, project_id: str, pay
user_id=user_id, user_id=user_id,
project_id=project.id, project_id=project.id,
message="更新真人素材项目", message="更新真人素材项目",
detail={"before": before, "after": after, "remote_name_locked": remote_name_locked}, detail={"before": before, "after": after},
) )
return project return project