227 lines
9.3 KiB
Python
227 lines
9.3 KiB
Python
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)
|