fix(account): 支持跨平台同名账号——(platform,name)联合唯一+Asset.account_id外键重构
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
"""平台账号业务逻辑
|
||||
|
||||
账号与资产的关系:Asset.account 按名称引用账号(字符串,兼容历史自由文本)。
|
||||
重命名账号时同步更新所有引用资产,保证两边一致。
|
||||
账号与资产的关系:Asset.account_id 外键关联账号(重命名账号不影响引用)。
|
||||
唯一性:(platform, name) 联合唯一——同一邮箱/用户名可跨平台复用,同平台内不重名。
|
||||
凭证层:登录密码与 API 配置加密存储,Read 仅返回布尔标记。
|
||||
"""
|
||||
|
||||
@@ -16,19 +16,19 @@ from app.models.asset import Account, Asset
|
||||
from app.schemas.account import AccountCreate, AccountRead, AccountUpdate
|
||||
|
||||
|
||||
def _asset_counts(session: Session) -> Dict[str, int]:
|
||||
"""按 account 名称统计引用资产数"""
|
||||
def _asset_counts(session: Session) -> Dict[int, int]:
|
||||
"""按 account_id 统计引用资产数"""
|
||||
rows = session.exec(
|
||||
select(Asset.account, func.count(Asset.id))
|
||||
.where(Asset.account.is_not(None)) # type: ignore[union-attr]
|
||||
.group_by(Asset.account)
|
||||
select(Asset.account_id, func.count(Asset.id))
|
||||
.where(Asset.account_id.is_not(None)) # type: ignore[union-attr]
|
||||
.group_by(Asset.account_id)
|
||||
).all()
|
||||
return {name: cnt for name, cnt in rows}
|
||||
return {aid: cnt for aid, cnt in rows}
|
||||
|
||||
|
||||
def _to_read(account: Account, counts: Dict[str, int]) -> AccountRead:
|
||||
def _to_read(account: Account, counts: Dict[int, int]) -> AccountRead:
|
||||
read = AccountRead.model_validate(account)
|
||||
read.asset_count = counts.get(account.name, 0)
|
||||
read.asset_count = counts.get(account.id, 0)
|
||||
read.has_login_password = bool(account.login_password_encrypted)
|
||||
read.has_api_config = bool(account.api_config_encrypted)
|
||||
return read
|
||||
@@ -47,13 +47,20 @@ def _get_account(session: Session, account_id: int) -> Account:
|
||||
return account
|
||||
|
||||
|
||||
def _check_name_taken(session: Session, name: str, exclude_id: int | None = None) -> None:
|
||||
stmt = select(Account).where(Account.name == name)
|
||||
def _norm_platform(platform: str | None) -> str:
|
||||
"""platform 规范化为非空字符串(避免 NULL 绕过 (platform,name) 唯一约束)"""
|
||||
return (platform or "").strip()
|
||||
|
||||
|
||||
def _check_name_taken(session: Session, name: str, platform: str, exclude_id: int | None = None) -> None:
|
||||
"""校验 (platform, name) 联合唯一:同平台内不允许重名,跨平台可复用"""
|
||||
stmt = select(Account).where(Account.name == name, Account.platform == platform)
|
||||
if exclude_id is not None:
|
||||
stmt = stmt.where(Account.id != exclude_id)
|
||||
if session.exec(stmt).first():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail=f"账号已存在:{name}"
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"该平台下账号已存在:{name}",
|
||||
)
|
||||
|
||||
|
||||
@@ -61,8 +68,9 @@ def create_account(session: Session, data: AccountCreate) -> AccountRead:
|
||||
name = data.name.strip()
|
||||
if not name:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="账号名称不能为空")
|
||||
_check_name_taken(session, name)
|
||||
account = Account(name=name, platform=data.platform, remark=data.remark, login_user=data.login_user)
|
||||
platform = _norm_platform(data.platform)
|
||||
_check_name_taken(session, name, platform)
|
||||
account = Account(name=name, platform=platform, remark=data.remark, login_user=data.login_user)
|
||||
account.login_password_encrypted = crypto.encrypt(data.login_password)
|
||||
account.api_config_encrypted = crypto.encrypt(data.api_config)
|
||||
session.add(account)
|
||||
@@ -73,20 +81,17 @@ def create_account(session: Session, data: AccountCreate) -> AccountRead:
|
||||
|
||||
def update_account(session: Session, account_id: int, data: AccountUpdate) -> AccountRead:
|
||||
account = _get_account(session, account_id)
|
||||
# platform 可能随本次更新变化,校验重名时用更新后的值
|
||||
new_platform = _norm_platform(data.platform) if data.platform is not None else _norm_platform(account.platform)
|
||||
if data.name is not None:
|
||||
new_name = data.name.strip()
|
||||
if not new_name:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="账号名称不能为空")
|
||||
if new_name != account.name:
|
||||
_check_name_taken(session, new_name, exclude_id=account.id)
|
||||
# 同步更新引用该账号的资产,避免重命名后资产端失联
|
||||
assets = session.exec(select(Asset).where(Asset.account == account.name)).all()
|
||||
for a in assets:
|
||||
a.account = new_name
|
||||
session.add(a)
|
||||
if new_name != account.name or new_platform != _norm_platform(account.platform):
|
||||
_check_name_taken(session, new_name, new_platform, exclude_id=account.id)
|
||||
account.name = new_name
|
||||
if data.platform is not None:
|
||||
account.platform = data.platform or None
|
||||
account.platform = new_platform
|
||||
if data.remark is not None:
|
||||
account.remark = data.remark or None
|
||||
if data.login_user is not None:
|
||||
@@ -104,10 +109,14 @@ def update_account(session: Session, account_id: int, data: AccountUpdate) -> Ac
|
||||
|
||||
def delete_account(session: Session, account_id: int) -> Dict[str, int]:
|
||||
account = _get_account(session, account_id)
|
||||
affected = len(session.exec(select(Asset.id).where(Asset.account == account.name)).all())
|
||||
# 引用该账号的资产:account_id 置 NULL(保留资产,仅解除关联)
|
||||
assets = session.exec(select(Asset).where(Asset.account_id == account_id)).all()
|
||||
for a in assets:
|
||||
a.account_id = None
|
||||
session.add(a)
|
||||
affected = len(assets)
|
||||
session.delete(account)
|
||||
session.commit()
|
||||
# 资产端保留原账号名文本(不级联清空),由用户自行处理
|
||||
return {"affected_assets": affected}
|
||||
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from sqlmodel import Session, select
|
||||
|
||||
from app.core import crypto
|
||||
from app.models.asset import (
|
||||
Account,
|
||||
AIAccount,
|
||||
Asset,
|
||||
AssetStatus,
|
||||
@@ -93,11 +94,29 @@ def _provider_name(session: Session, asset: Asset) -> Optional[str]:
|
||||
return None
|
||||
|
||||
|
||||
def _account_name(session: Session, asset: Asset) -> Optional[str]:
|
||||
"""取关联账号标识(展示用)"""
|
||||
if asset.account_id:
|
||||
account = session.get(Account, asset.account_id)
|
||||
if account:
|
||||
return account.name
|
||||
return None
|
||||
|
||||
|
||||
def _validate_account(session: Session, account_id: Optional[int]) -> None:
|
||||
"""account_id 非空时确认账号存在,防外键悬空"""
|
||||
if account_id is not None and not session.get(Account, account_id):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail=f"账号不存在:id={account_id}"
|
||||
)
|
||||
|
||||
|
||||
def _to_read(session: Session, asset: Asset, detail) -> AssetRead:
|
||||
"""组装 AssetRead 输出(主表 + 详情 + 计算字段)"""
|
||||
read = AssetRead.model_validate(asset)
|
||||
read.days_to_expiry = _days_to_expiry(asset.expiry_date)
|
||||
read.provider_name = _provider_name(session, asset)
|
||||
read.account_name = _account_name(session, asset)
|
||||
if isinstance(detail, VPSDetail):
|
||||
read.vps_detail = _detail_to_read(detail)
|
||||
elif isinstance(detail, DomainDetail):
|
||||
@@ -156,6 +175,12 @@ def _get_details_batch(session: Session, assets: list) -> dict:
|
||||
if provider_ids:
|
||||
for p in session.exec(select(Provider).where(Provider.id.in_(provider_ids))).all():
|
||||
provider_map[p.id] = p.name
|
||||
# 批量查 Account 名
|
||||
account_ids = {a.account_id for a in assets if a.account_id}
|
||||
account_map = {}
|
||||
if account_ids:
|
||||
for acc in session.exec(select(Account).where(Account.id.in_(account_ids))).all():
|
||||
account_map[acc.id] = acc.name
|
||||
# 批量查各类型详情
|
||||
detail_map = {}
|
||||
for asset_type, (_, model) in DETAIL_MAP.items():
|
||||
@@ -163,7 +188,7 @@ def _get_details_batch(session: Session, assets: list) -> dict:
|
||||
if typed_ids:
|
||||
for d in session.exec(select(model).where(model.asset_id.in_(typed_ids))).all():
|
||||
detail_map[d.asset_id] = d
|
||||
return {"providers": provider_map, "details": detail_map}
|
||||
return {"providers": provider_map, "accounts": account_map, "details": detail_map}
|
||||
|
||||
|
||||
def _to_read_batch(session: Session, asset: Asset, batch: dict) -> AssetRead:
|
||||
@@ -171,6 +196,7 @@ def _to_read_batch(session: Session, asset: Asset, batch: dict) -> AssetRead:
|
||||
read = AssetRead.model_validate(asset)
|
||||
read.days_to_expiry = _days_to_expiry(asset.expiry_date)
|
||||
read.provider_name = batch["providers"].get(asset.provider_id)
|
||||
read.account_name = batch.get("accounts", {}).get(asset.account_id)
|
||||
detail = batch["details"].get(asset.id)
|
||||
if isinstance(detail, VPSDetail):
|
||||
read.vps_detail = _detail_to_read(detail)
|
||||
@@ -203,6 +229,7 @@ def _create_asset_no_commit(session: Session, data: AssetCreate) -> Asset:
|
||||
asset_data = data.model_dump(
|
||||
exclude={"vps_detail", "domain_detail", "ai_detail", "cloudflare_detail"}
|
||||
)
|
||||
_validate_account(session, asset_data.get("account_id"))
|
||||
asset = Asset(**asset_data)
|
||||
session.add(asset)
|
||||
session.flush() # 获取 asset.id,但不提交
|
||||
@@ -239,6 +266,8 @@ def _update_asset_no_commit(session: Session, asset_id: int, data: AssetUpdate)
|
||||
main_fields = data.model_dump(
|
||||
exclude_unset=True, exclude={"vps_detail", "domain_detail", "ai_detail", "cloudflare_detail"}
|
||||
)
|
||||
if "account_id" in main_fields:
|
||||
_validate_account(session, main_fields["account_id"])
|
||||
for key, value in main_fields.items():
|
||||
setattr(asset, key, value)
|
||||
session.add(asset)
|
||||
|
||||
@@ -169,7 +169,7 @@ def _find_asset(session: Session, provider_id: int, external_id: str, asset_type
|
||||
).first()
|
||||
|
||||
|
||||
def _sync_vps(session: Session, provider: Provider, adapter: BaseAdapter, account_name: Optional[str] = None) -> dict:
|
||||
def _sync_vps(session: Session, provider: Provider, adapter: BaseAdapter, account_id: Optional[int] = None) -> dict:
|
||||
created = updated = 0
|
||||
for vps in adapter.list_vps():
|
||||
existing = _find_asset(session, provider.id, vps.external_id, AssetType.VPS)
|
||||
@@ -179,8 +179,8 @@ def _sync_vps(session: Session, provider: Provider, adapter: BaseAdapter, accoun
|
||||
if vps.monthly_cost is not None:
|
||||
existing.cost = vps.monthly_cost
|
||||
existing.currency = vps.currency
|
||||
if account_name:
|
||||
existing.account = account_name
|
||||
if account_id:
|
||||
existing.account_id = account_id
|
||||
session.add(existing)
|
||||
detail = session.exec(
|
||||
select(VPSDetail).where(VPSDetail.asset_id == existing.id)
|
||||
@@ -201,7 +201,7 @@ def _sync_vps(session: Session, provider: Provider, adapter: BaseAdapter, accoun
|
||||
provider=provider.slug,
|
||||
provider_id=provider.id,
|
||||
external_id=vps.external_id,
|
||||
account=account_name,
|
||||
account_id=account_id,
|
||||
status=_norm_status(vps.status),
|
||||
cost=vps.monthly_cost or 0,
|
||||
currency=vps.currency,
|
||||
@@ -224,7 +224,7 @@ def _sync_vps(session: Session, provider: Provider, adapter: BaseAdapter, accoun
|
||||
return {"created": created, "updated": updated}
|
||||
|
||||
|
||||
def _sync_domains(session: Session, provider: Provider, adapter: BaseAdapter, account_name: Optional[str] = None) -> dict:
|
||||
def _sync_domains(session: Session, provider: Provider, adapter: BaseAdapter, account_id: Optional[int] = None) -> dict:
|
||||
created = updated = 0
|
||||
for dom in adapter.list_domains():
|
||||
existing = _find_asset(session, provider.id, dom.external_id, AssetType.DOMAIN)
|
||||
@@ -232,8 +232,8 @@ def _sync_domains(session: Session, provider: Provider, adapter: BaseAdapter, ac
|
||||
existing.status = _norm_status(dom.status)
|
||||
if dom.expiry_date:
|
||||
existing.expiry_date = dom.expiry_date
|
||||
if account_name:
|
||||
existing.account = account_name
|
||||
if account_id:
|
||||
existing.account_id = account_id
|
||||
session.add(existing)
|
||||
detail = session.exec(
|
||||
select(DomainDetail).where(DomainDetail.asset_id == existing.id)
|
||||
@@ -250,7 +250,7 @@ def _sync_domains(session: Session, provider: Provider, adapter: BaseAdapter, ac
|
||||
provider=provider.slug,
|
||||
provider_id=provider.id,
|
||||
external_id=dom.external_id,
|
||||
account=account_name,
|
||||
account_id=account_id,
|
||||
status=_norm_status(dom.status),
|
||||
expiry_date=dom.expiry_date,
|
||||
)
|
||||
@@ -295,12 +295,12 @@ def sync_account(session: Session, account_id: int) -> dict:
|
||||
|
||||
if caps["list_vps"]:
|
||||
try:
|
||||
result["vps"] = _sync_vps(session, provider, adapter, account.name)
|
||||
result["vps"] = _sync_vps(session, provider, adapter, account.id)
|
||||
except Exception as e: # noqa: BLE001
|
||||
result["vps_error"] = str(e)
|
||||
if caps["list_domains"]:
|
||||
try:
|
||||
result["domains"] = _sync_domains(session, provider, adapter, account.name)
|
||||
result["domains"] = _sync_domains(session, provider, adapter, account.id)
|
||||
except Exception as e: # noqa: BLE001
|
||||
result["domains_error"] = str(e)
|
||||
if caps["get_account"]:
|
||||
|
||||
Reference in New Issue
Block a user