Files

188 lines
7.1 KiB
Python

"""火山引擎人像素材库 Assets API 客户端(AK/SK 签名)。
照搬自 AirDrama,改用 settings.ASSETS_API 凭证(暂借 AirDrama 已邀测账号,张业昌待换)。
同步调用,出错抛 AssetsAPIError。真人(person)基础资产生成后静默上传审核 → 轮询 active(绿)/failed(红)。
"""
import json
import logging
from django.conf import settings
logger = logging.getLogger(__name__)
SERVICE = "ark"
REGION = "cn-beijing"
API_VERSION = "2024-01-01"
HOST = "open.volcengineapi.com"
_ASSETS_ERROR_MESSAGES = {
"ConfigError": "素材审核服务未配置(ASSETS_API_* 凭证)",
"RequestError": "素材审核服务暂时不可用,请稍后重试",
"InvalidParameter": "素材参数无效",
"NotFound": "素材不存在或已被删除",
"NotFound.asset_id": "素材不存在或已被删除",
"NotFound.group_id": "素材组不存在或已被删除",
"Forbidden": "没有权限操作该素材(检查 AK/SK)",
}
_NOT_FOUND_CODES = frozenset({"NotFound", "NotFound.asset_id", "NotFound.group_id"})
_ASSET_NOT_FOUND_CODES = frozenset({"NotFound", "NotFound.asset_id"})
_GROUP_NOT_FOUND_CODES = frozenset({"NotFound", "NotFound.group_id"})
class AssetsAPIError(Exception):
def __init__(self, code, message, status_code=400):
self.code = code
self.api_message = message
self.status_code = status_code
self.user_message = _ASSETS_ERROR_MESSAGES.get(code) or "素材审核操作失败,请稍后重试"
super().__init__(f"[{code}] {message}")
def is_not_found_code(code: str | None) -> bool:
return code in _NOT_FOUND_CODES
def is_asset_not_found_code(code: str | None) -> bool:
return code in _ASSET_NOT_FOUND_CODES
def is_group_not_found_code(code: str | None) -> bool:
return code in _GROUP_NOT_FOUND_CODES
def is_enabled() -> bool:
cfg = getattr(settings, "ASSETS_API", {}) or {}
return bool(cfg.get("enabled") and cfg.get("access_key") and cfg.get("secret_key"))
def _project() -> str:
return (getattr(settings, "ASSETS_API", {}) or {}).get("project_name") or "int_dev_Airlabs"
def _get_service():
"""构建带 AK/SK 的 volcengine Service(延迟 import SDK,未装/未配不影响其它功能)。"""
cfg = getattr(settings, "ASSETS_API", {}) or {}
ak, sk = cfg.get("access_key"), cfg.get("secret_key")
if not ak or not sk:
raise AssetsAPIError("ConfigError", "ASSETS_API access_key / secret_key not configured")
from volcengine.ApiInfo import ApiInfo
from volcengine.base.Service import Service
from volcengine.Credentials import Credentials
from volcengine.ServiceInfo import ServiceInfo
service_info = ServiceInfo(
HOST,
{"Accept": "application/json", "Content-Type": "application/json"},
Credentials(ak, sk, SERVICE, REGION),
10, 30,
)
actions = [
"CreateAssetGroup", "CreateAsset", "ListAssetGroups", "ListAssets", "GetAsset", "DeleteAsset",
# 自由创作素材库(组/素材全生命周期)
"GetAssetGroup", "UpdateAssetGroup", "UpdateAsset", "DeleteAssetGroup",
]
api_info = {a: ApiInfo("POST", "/", {"Action": a, "Version": API_VERSION}, {}, {}) for a in actions}
return Service(service_info, api_info)
def _do_request(action: str, body_dict: dict) -> dict:
service = _get_service()
body = json.dumps(body_dict, ensure_ascii=False)
try:
resp = service.json(action, {}, body)
except Exception as e: # noqa: BLE001 — SDK 非 200 抛 Exception(resp.text.encode())
raw = e.args[0] if e.args else ""
error_str = raw.decode("utf-8") if isinstance(raw, bytes) else str(raw)
logger.warning("Assets API %s error: %s", action, error_str[:300])
try:
err = json.loads(error_str).get("ResponseMetadata", {}).get("Error", {})
if err:
raise AssetsAPIError(err.get("Code", "Unknown"), err.get("Message", error_str))
except json.JSONDecodeError:
pass
raise AssetsAPIError("RequestError", error_str or "Empty response")
data = json.loads(resp) if isinstance(resp, str) else resp
err = data.get("ResponseMetadata", {}).get("Error", {})
if err:
raise AssetsAPIError(err.get("Code", "Unknown"), err.get("Message", str(data)))
return data.get("Result", {})
def create_asset_group(name: str, description: str = "", group_type: str = "AIGC") -> str:
"""建素材组,返回远程 group id。"""
return _do_request(
"CreateAssetGroup",
{"Name": name, "Description": description, "GroupType": group_type, "ProjectName": _project()},
).get("Id", "")
def create_asset(group_id: str, image_url: str, name: str = "", asset_type: str = "Image") -> str:
"""往组里传一张素材(URL),返回远程 asset id;审核异步,稍后 get_asset 查状态。"""
return _do_request(
"CreateAsset",
{"GroupId": group_id, "URL": image_url, "Name": name, "AssetType": asset_type, "ProjectName": _project()},
).get("Id", "")
def get_asset(asset_id: str) -> dict:
"""查单个素材详情(含审核 Status:Active/Failed/Processing、Url、ErrorMessage)。"""
return _do_request("GetAsset", {"Id": asset_id, "ProjectName": _project()})
def list_asset_groups(page: int = 1, page_size: int = 20, name: str | None = None) -> tuple:
filter_dict = {"GroupType": "AIGC"}
if name:
filter_dict["Name"] = name
result = _do_request(
"ListAssetGroups",
{"Filter": filter_dict, "PageNumber": page, "PageSize": page_size, "ProjectName": _project()},
)
return result.get("Items", []), result.get("TotalCount", 0)
def list_assets(group_ids: list | None = None, status: str | None = None,
name: str | None = None, page: int = 1, page_size: int = 20) -> tuple:
"""列组内素材。返回 (items, total_count)。"""
filter_dict: dict = {"GroupType": "AIGC"}
if group_ids:
filter_dict["GroupIds"] = group_ids
if status:
filter_dict["Statuses"] = [status]
if name:
filter_dict["Name"] = name
result = _do_request(
"ListAssets",
{"Filter": filter_dict, "PageNumber": page, "PageSize": page_size, "ProjectName": _project()},
)
return result.get("Items", []), result.get("TotalCount", 0)
def get_asset_group(group_id: str) -> dict:
return _do_request("GetAssetGroup", {"Id": group_id, "ProjectName": _project()})
def update_asset_group(group_id: str, name: str | None = None, description: str | None = None) -> None:
body: dict = {"Id": group_id, "ProjectName": _project()}
if name is not None:
body["Name"] = name
if description is not None:
body["Description"] = description
_do_request("UpdateAssetGroup", body)
def update_asset(asset_id: str, name: str | None = None) -> None:
body: dict = {"Id": asset_id, "ProjectName": _project()}
if name is not None:
body["Name"] = name
_do_request("UpdateAsset", body)
def delete_asset(asset_id: str) -> None:
_do_request("DeleteAsset", {"Id": asset_id, "ProjectName": _project()})
def delete_asset_group(group_id: str) -> None:
"""删组(远程级联删组内素材)。"""
_do_request("DeleteAssetGroup", {"Id": group_id, "ProjectName": _project()})