from __future__ import annotations import asyncio import json from typing import Any from fastapi import HTTPException 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 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: """火山 Ark 私域可信素材 Asset API Client。只做 AK/SK 鉴权调用与响应标准化。""" def __init__(self, *, for_celery: bool = False): self.ak = settings.VOLC_ACCESS_KEY_ID self.sk = settings.VOLC_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_group(self, *, project_name: str, name: str, description: str | None = None, group_type: str = "AIGC") -> dict[str, Any]: payload: dict[str, Any] = {"Name": name, "GroupType": group_type, "ProjectName": project_name} if description: payload["Description"] = description return await self._call(ArkPrivatePortraitAction.CREATE_ASSET_GROUP, payload) 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 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: 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), ) 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]: 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, "https", ) # volcengine SDK 的 ApiInfo.query 必须是 dict,不能传 "Action=xxx&Version=xxx" 字符串。 # volcengine SDK 签名时要求 body 是 bytes。 # ServiceInfo 默认 scheme='http',这里必须显式传 https,避免请求走 http://host:80。 api_info = { action.value: ApiInfo( "POST", "/", {"Action": action.value, "Version": ARK_PRIVATE_PORTRAIT_VERSION}, {}, {}, ) } service = Service(service_info, api_info) body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8") try: raw = service.json(action.value, {}, body) 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) if "ConnectTimeout" in message or "timed out" in message or "Connection" in message: raise ArkPrivateAssetClientError(f"{action.value} 连接火山私域素材接口失败:{message}") from exc raise ArkPrivateAssetClientError(f"{action.value} 请求火山私域素材接口异常:{message}") from exc resp = self._normalize_response(raw) self._raise_remote_error_if_any(action.value, resp) metadata = resp.get("ResponseMetadata") if isinstance(resp, dict) else None request_id = metadata.get("RequestId") if isinstance(metadata, dict) else None 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} @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 def _exception_message(exc: Exception) -> str: if exc.args: raw = exc.args[0] if isinstance(raw, (bytes, bytearray)): return raw.decode("utf-8", errors="ignore") return str(raw) return str(exc) @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}