180 lines
9.0 KiB
Python
180 lines
9.0 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
from typing import Any
|
|
|
|
from app.config import settings
|
|
from app.enums.private_portrait import (
|
|
ARK_PRIVATE_PORTRAIT_HOST,
|
|
ARK_PRIVATE_PORTRAIT_REGION,
|
|
ARK_PRIVATE_PORTRAIT_SERVICE_NAME,
|
|
ARK_PRIVATE_PORTRAIT_VERSION,
|
|
ArkPrivatePortraitAction,
|
|
PrivatePortraitEventSource,
|
|
PrivatePortraitEventStatus,
|
|
PrivatePortraitEventType,
|
|
)
|
|
from app.services.operation_log_service import log_remote_api_event
|
|
from app.services.private_portrait.rate_limiter import acquire_private_portrait_action_token
|
|
|
|
DOMAIN = "private_portrait"
|
|
|
|
|
|
class ArkPrivateAssetClientError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class ArkPrivateAssetClient:
|
|
"""火山 Ark 私域真人人像素材 API Client。只做 AK/SK 鉴权调用与响应标准化。"""
|
|
|
|
def __init__(self, *, ak: str | None = None, sk: str | None = None, for_celery: bool = False):
|
|
self.ak = ak or settings.VOLC_SMS_ACCESS_KEY_ID
|
|
self.sk = sk or settings.VOLC_SMS_SECRET_ACCESS_KEY
|
|
self.for_celery = for_celery
|
|
if not self.ak or not self.sk:
|
|
raise ArkPrivateAssetClientError("火山 AK/SK 未配置:VOLC_SMS_ACCESS_KEY_ID / VOLC_SMS_SECRET_ACCESS_KEY")
|
|
|
|
async def create_visual_validate_session(self, *, project_name: str, callback_url: str) -> dict[str, Any]:
|
|
return await self._call(ArkPrivatePortraitAction.CREATE_VISUAL_VALIDATE_SESSION, {"CallbackURL": callback_url, "ProjectName": project_name})
|
|
|
|
async def get_visual_validate_result(self, *, project_name: str, byted_token: str) -> dict[str, Any]:
|
|
return await self._call(ArkPrivatePortraitAction.GET_VISUAL_VALIDATE_RESULT, {"BytedToken": byted_token, "ProjectName": project_name})
|
|
|
|
async def create_asset(self, *, project_name: str, group_id: str, url: str, asset_type: str, name: str | None = None) -> dict[str, Any]:
|
|
payload: dict[str, Any] = {"GroupId": group_id, "URL": url, "AssetType": asset_type, "ProjectName": project_name}
|
|
if name:
|
|
payload["Name"] = name
|
|
return await self._call(ArkPrivatePortraitAction.CREATE_ASSET, payload)
|
|
|
|
async def get_asset(self, *, project_name: str, asset_id: str) -> dict[str, Any]:
|
|
return await self._call(ArkPrivatePortraitAction.GET_ASSET, {"Id": asset_id, "ProjectName": project_name})
|
|
|
|
async def list_assets(self, *, project_name: str, filter_payload: dict[str, Any] | None = None, page_number: int = 1, page_size: int = 20) -> dict[str, Any]:
|
|
payload = {"Filter": filter_payload or {}, "PageNumber": page_number, "PageSize": page_size, "ProjectName": project_name}
|
|
return await self._call(ArkPrivatePortraitAction.LIST_ASSETS, payload)
|
|
|
|
async def list_asset_groups(self, *, project_name: str, filter_payload: dict[str, Any] | None = None, page_number: int = 1, page_size: int = 20) -> dict[str, Any]:
|
|
payload = {"Filter": filter_payload or {}, "PageNumber": page_number, "PageSize": page_size, "ProjectName": project_name}
|
|
return await self._call(ArkPrivatePortraitAction.LIST_ASSET_GROUPS, payload)
|
|
|
|
async def get_asset_group(self, *, project_name: str, group_id: str) -> dict[str, Any]:
|
|
return await self._call(ArkPrivatePortraitAction.GET_ASSET_GROUP, {"Id": group_id, "ProjectName": project_name})
|
|
|
|
async def update_asset_group(self, *, project_name: str, group_id: str, name: str | None = None, title: str | None = None, description: str | None = None) -> dict[str, Any]:
|
|
payload: dict[str, Any] = {"Id": group_id, "ProjectName": project_name}
|
|
if name is not None:
|
|
payload["Name"] = name
|
|
if title is not None:
|
|
payload["Title"] = title
|
|
if description is not None:
|
|
payload["Description"] = description
|
|
return await self._call(ArkPrivatePortraitAction.UPDATE_ASSET_GROUP, payload)
|
|
|
|
async def update_asset(self, *, project_name: str, asset_id: str, name: str | None = None) -> dict[str, Any]:
|
|
payload: dict[str, Any] = {"Id": asset_id, "ProjectName": project_name}
|
|
if name is not None:
|
|
payload["Name"] = name
|
|
return await self._call(ArkPrivatePortraitAction.UPDATE_ASSET, payload)
|
|
|
|
async def delete_asset(self, *, project_name: str, asset_id: str) -> dict[str, Any]:
|
|
return await self._call(ArkPrivatePortraitAction.DELETE_ASSET, {"Id": asset_id, "ProjectName": project_name})
|
|
|
|
async def delete_asset_group(self, *, project_name: str, group_id: str) -> dict[str, Any]:
|
|
return await self._call(ArkPrivatePortraitAction.DELETE_ASSET_GROUP, {"Id": group_id, "ProjectName": project_name})
|
|
|
|
async def _call(self, action: ArkPrivatePortraitAction, payload: dict[str, Any]) -> dict[str, Any]:
|
|
action_value = action.value
|
|
await acquire_private_portrait_action_token(action=action_value, wait_timeout_seconds=2.0, for_celery=self.for_celery)
|
|
log_remote_api_event(
|
|
domain=DOMAIN,
|
|
remote_action=action_value,
|
|
event_type=PrivatePortraitEventType.ARK_API_CALL_START.value,
|
|
event_status=PrivatePortraitEventStatus.PENDING.value,
|
|
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
|
|
request=payload,
|
|
)
|
|
try:
|
|
result = await asyncio.to_thread(self._call_sync, action, payload)
|
|
log_remote_api_event(
|
|
domain=DOMAIN,
|
|
remote_action=action_value,
|
|
event_type=PrivatePortraitEventType.ARK_API_CALL_SUCCESS.value,
|
|
event_status=PrivatePortraitEventStatus.SUCCESS.value,
|
|
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
|
|
request=payload,
|
|
response=result,
|
|
remote_request_id=result.get("RequestId") or result.get("request_id"),
|
|
)
|
|
return result
|
|
except Exception 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,
|
|
remote_message=str(exc),
|
|
)
|
|
raise
|
|
|
|
def _call_sync(self, action: ArkPrivatePortraitAction, payload: dict[str, Any]) -> dict[str, Any]:
|
|
try:
|
|
from volcengine.ApiInfo import ApiInfo
|
|
from volcengine.Credentials import Credentials
|
|
from volcengine.ServiceInfo import ServiceInfo
|
|
from volcengine.base.Service import Service
|
|
except Exception as exc:
|
|
raise ArkPrivateAssetClientError("缺少火山 volcengine Python SDK。请确认线上环境已安装 volcengine。") from exc
|
|
|
|
credentials = Credentials(self.ak, self.sk, ARK_PRIVATE_PORTRAIT_SERVICE_NAME, ARK_PRIVATE_PORTRAIT_REGION)
|
|
service_info = ServiceInfo(
|
|
ARK_PRIVATE_PORTRAIT_HOST,
|
|
{"Content-Type": "application/json", "Accept": "application/json"},
|
|
credentials,
|
|
10,
|
|
60,
|
|
)
|
|
api_info = {action.value: ApiInfo("POST", "/", f"Action={action.value}&Version={ARK_PRIVATE_PORTRAIT_VERSION}", {}, {})}
|
|
service = Service(service_info, api_info)
|
|
try:
|
|
raw = service.json(action.value, {}, json.dumps(payload, ensure_ascii=False, separators=(",", ":")))
|
|
except TypeError:
|
|
raw = service.json(action.value, {}, payload)
|
|
|
|
resp = self._normalize_response(raw)
|
|
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
|
|
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):
|
|
result = dict(resp["Result"])
|
|
if request_id:
|
|
result["RequestId"] = request_id
|
|
return result
|
|
if isinstance(resp, dict):
|
|
if request_id:
|
|
resp.setdefault("RequestId", request_id)
|
|
return resp
|
|
return {"raw": resp, "RequestId": request_id}
|
|
|
|
@staticmethod
|
|
def _normalize_response(raw: Any) -> dict[str, Any]:
|
|
if raw is None:
|
|
return {}
|
|
if isinstance(raw, dict):
|
|
return raw
|
|
if isinstance(raw, (bytes, bytearray)):
|
|
raw = raw.decode("utf-8", errors="ignore")
|
|
if isinstance(raw, str):
|
|
try:
|
|
obj = json.loads(raw)
|
|
return obj if isinstance(obj, dict) else {"raw": obj}
|
|
except json.JSONDecodeError:
|
|
return {"raw": raw}
|
|
return {"raw": raw}
|