from __future__ import annotations from collections.abc import Iterable from typing import Any from app.accounts.repository import ( AWS_SECRET_NAMES, AccountProvider, AccountRepository, CredentialAccountRecord, ) from app.accounts.schemas import ( AccountCreate, AccountSecretFlags, AccountUpdate, AccountView, ) from app.core.errors import ConflictError, NotFoundError, ValidationAppError from app.integrations.repository import IntegrationRepository class AccountService: def __init__(self, repository: AccountRepository) -> None: self.repository = repository def list_accounts(self, provider: AccountProvider | None = None) -> list[dict[str, Any]]: accounts = self.repository.list_accounts(provider) flags = self.repository.secret_flags_many(account.id for account in accounts) return [self._view(account, flags[account.id]).model_dump() for account in accounts] def get_account(self, account_id: str) -> dict[str, Any]: account = self._require_account(account_id) return self._view(account, self.repository.secret_flags(account.id)).model_dump() def create_account(self, payload: AccountCreate) -> dict[str, Any]: if payload.provider == "aws": use_default: bool | None = payload.use_default_aws_credentials secrets = self._provided_secrets(payload, AWS_SECRET_NAMES) else: use_default = None secrets = self._provided_secrets(payload, ("cloudflare_api_token",)) try: account = self.repository.create_account( provider=payload.provider, name=payload.name, use_default_aws_credentials=use_default, secret_values=secrets, ) except RuntimeError as exc: raise self._repository_error(exc) from exc return self.get_account(account.id) def import_legacy_accounts( self, integration_repository: IntegrationRepository, ) -> dict[AccountProvider, str]: settings = integration_repository.get() legacy_secrets = integration_repository.get_secrets() access_key = legacy_secrets.get("aws_access_key_id") secret_key = legacy_secrets.get("aws_secret_access_key") aws_configured = settings.use_default_aws_credentials or bool(access_key and secret_key) aws_secrets: dict[str, str] = {} if access_key and secret_key: aws_secrets = { "aws_access_key_id": access_key, "aws_secret_access_key": secret_key, } if session_token := legacy_secrets.get("aws_session_token"): aws_secrets["aws_session_token"] = session_token return self.repository.import_legacy_accounts( aws_use_default=(settings.use_default_aws_credentials if aws_configured else None), aws_secret_values=aws_secrets, cloudflare_api_token=legacy_secrets.get("cloudflare_api_token"), ) def update_account(self, account_id: str, payload: AccountUpdate) -> dict[str, Any]: account = self._require_account(account_id) values: dict[str, object] = {} if "name" in payload.model_fields_set: if payload.name is None: raise ValidationAppError("账号名称不能为空") values["name"] = payload.name secret_values: dict[str, str] = {} delete_names: set[str] = set() if account.provider == "aws": self._reject_cloudflare_fields(payload) self._prepare_aws_update( account, payload, values, secret_values, delete_names, ) else: self._reject_aws_fields(payload) if "cloudflare_api_token" in payload.model_fields_set and payload.cloudflare_api_token: secret_values["cloudflare_api_token"] = payload.cloudflare_api_token try: updated = self.repository.update_account( account_id, expected_version=payload.config_version, values=values, secret_values=secret_values, delete_secret_names=delete_names, ) except RuntimeError as exc: raise self._repository_error(exc) from exc return self.get_account(updated.id) def archive_account(self, account_id: str) -> dict[str, Any]: try: account = self.repository.archive_account(account_id) except RuntimeError as exc: raise self._repository_error(exc) from exc return self._view(account, self.repository.secret_flags(account.id)).model_dump() def resolve_aws(self, account_id: str) -> tuple[bool, dict[str, str | None]]: try: return self.repository.resolve_aws(account_id) except RuntimeError as exc: raise self._repository_error(exc) from exc def resolve_cloudflare(self, account_id: str) -> str: try: return self.repository.resolve_cloudflare(account_id) except RuntimeError as exc: raise self._repository_error(exc) from exc def _prepare_aws_update( self, account: CredentialAccountRecord, payload: AccountUpdate, values: dict[str, object], secret_values: dict[str, str], delete_names: set[str], ) -> None: fields = payload.model_fields_set if "use_default_aws_credentials" in fields: if payload.use_default_aws_credentials is None: raise ValidationAppError("请选择 AWS 凭据模式") values["use_default_aws_credentials"] = payload.use_default_aws_credentials final_use_default = bool( values.get( "use_default_aws_credentials", account.use_default_aws_credentials, ) ) pair_supplied = bool(payload.aws_access_key_id and payload.aws_secret_access_key) if pair_supplied: secret_values["aws_access_key_id"] = str(payload.aws_access_key_id) secret_values["aws_secret_access_key"] = str(payload.aws_secret_access_key) if payload.aws_session_token: secret_values["aws_session_token"] = payload.aws_session_token elif pair_supplied: delete_names.add("aws_session_token") if final_use_default: if pair_supplied or payload.aws_session_token: raise ValidationAppError("使用 AWS 默认凭据链时不能同时填写 API 密钥") delete_names.update(AWS_SECRET_NAMES) return flags = self.repository.secret_flags(account.id) has_pair = pair_supplied or (flags["aws_access_key_id"] and flags["aws_secret_access_key"]) if not has_pair: raise ValidationAppError("请填写 AWS API 密钥,或选择默认凭据链") @staticmethod def _reject_cloudflare_fields(payload: AccountUpdate) -> None: if "cloudflare_api_token" in payload.model_fields_set: raise ValidationAppError("AWS 账号不能包含 Cloudflare API Token") @staticmethod def _reject_aws_fields(payload: AccountUpdate) -> None: aws_fields = { "use_default_aws_credentials", "aws_access_key_id", "aws_secret_access_key", "aws_session_token", } if payload.model_fields_set & aws_fields: raise ValidationAppError("Cloudflare 账号不能包含 AWS 凭据") @staticmethod def _provided_secrets(payload: object, names: Iterable[str]) -> dict[str, str]: return { name: str(value) for name in names if (value := getattr(payload, name, None)) is not None } def _require_account(self, account_id: str) -> CredentialAccountRecord: account = self.repository.get_account(account_id) if account is None: raise NotFoundError("账号不存在", code="ACCOUNT_NOT_FOUND") return account @staticmethod def _view( account: CredentialAccountRecord, flags: dict[str, bool], ) -> AccountView: return AccountView( **account.to_dict(), secrets_configured=AccountSecretFlags(**flags), ) @staticmethod def _repository_error(exc: RuntimeError) -> Exception: code = str(exc) if code == "ACCOUNT_NOT_FOUND": return NotFoundError("账号不存在", code=code) conflicts = { "FLEET_RUN_ACTIVE": "轮换或 DNS 同步任务进行中,暂时不能修改账号", "ACCOUNT_NAME_CONFLICT": "同类型的账号名称已存在", "ACCOUNT_CONFIG_VERSION_CONFLICT": "账号配置已被其他操作更新,请刷新后重试", "ACCOUNT_IN_USE": "账号仍被活动实例使用,不能删除", } if code in conflicts: return ConflictError(conflicts[code], code=code) validations = { "ACCOUNT_PROVIDER_MISMATCH": "账号类型与所需凭据类型不匹配", "ACCOUNT_CREDENTIALS_INCOMPLETE": "账号凭据不完整,请先更新账号", "INVALID_ACCOUNT_DATA": "账号配置不合法", } if code in validations: return ValidationAppError(validations[code], code=code) return ValidationAppError("账号保存失败", code=code)